diff --git a/ms2deepscore/MS2DeepScoreEvaluated.py b/ms2deepscore/MS2DeepScoreEvaluated.py index 514704b0..a8c938a1 100644 --- a/ms2deepscore/MS2DeepScoreEvaluated.py +++ b/ms2deepscore/MS2DeepScoreEvaluated.py @@ -198,7 +198,6 @@ def _matrix_components( return result if MATCHMS_V1_API: - def matrix( self, spectra_1: List[Spectrum], @@ -219,7 +218,6 @@ def matrix( return as_matchms_scores(score_arrays) else: - def matrix( self, references: List[Spectrum], diff --git a/ms2deepscore/fingerprint_similarity_computations.py b/ms2deepscore/fingerprint_similarity_computations.py index bf450b26..010c3c7f 100644 --- a/ms2deepscore/fingerprint_similarity_computations.py +++ b/ms2deepscore/fingerprint_similarity_computations.py @@ -1,6 +1,6 @@ from typing import Tuple import numpy as np -from numba import jit, prange +from numba import jit, njit, prange from chemap.metrics import ( tanimoto_similarity_dense, tanimoto_similarity_sparse, @@ -84,6 +84,15 @@ def compute_fingerprint_similarity_row( )[0] +@njit +def _derive_seed(random_seed, row_index, bin_index): + return ( + random_seed + + 1_000_003 * row_index + + 97 * bin_index + ) % 4_294_967_295 + + # Add row based similarity computations # ------------------------------------- @@ -141,24 +150,50 @@ def _fill_pairs_for_row_same_set( max_pairs_per_bin, selection_bins, include_diagonal, + random_seed, ): + """Reservoir-sample at most ``max_pairs_per_bin`` candidates per bin. + + This avoids allocating/shuffling a potentially O(N) index array for every + score bin and row. Each qualifying pair has equal probability of ending up + in the bounded reservoir. + """ num_bins = len(selection_bins) for bin_number in range(num_bins): selection_bin = selection_bins[bin_number] - indices = np.nonzero( - (tanimoto_scores > selection_bin[0]) & (tanimoto_scores <= selection_bin[1]) - )[0] + # Derive an independent deterministic seed for each row/bin pair. + # This avoids dependence on prange/thread execution order. + if random_seed >= 0: + np.random.seed( + _derive_seed( + random_seed, + idx_fingerprint_i, + bin_number, + ) + ) - if not include_diagonal and idx_fingerprint_i in indices: - indices = indices[indices != idx_fingerprint_i] + seen = 0 + filled = 0 + for candidate_index in range(len(tanimoto_scores)): + if not include_diagonal and candidate_index == idx_fingerprint_i: + continue - np.random.shuffle(indices) - indices = indices[:max_pairs_per_bin] - num_indices = len(indices) + score = tanimoto_scores[candidate_index] + if not (score > selection_bin[0] and score <= selection_bin[1]): + continue - selected_pairs_per_bin[bin_number, idx_fingerprint_i, :num_indices] = indices - selected_scores_per_bin[bin_number, idx_fingerprint_i, :num_indices] = tanimoto_scores[indices] + seen += 1 + if filled < max_pairs_per_bin: + slot = filled + filled += 1 + else: + slot = np.random.randint(0, seen) + if slot >= max_pairs_per_bin: + continue + + selected_pairs_per_bin[bin_number, idx_fingerprint_i, slot] = candidate_index + selected_scores_per_bin[bin_number, idx_fingerprint_i, slot] = score @jit(nopython=True, parallel=True) @@ -167,6 +202,7 @@ def _compute_tanimoto_similarity_per_bin_dense( max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), include_diagonal=True, + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size = fingerprints.shape[0] num_bins = len(selection_bins) @@ -186,6 +222,7 @@ def _compute_tanimoto_similarity_per_bin_dense( max_pairs_per_bin, selection_bins, include_diagonal, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -197,6 +234,7 @@ def _compute_tanimoto_similarity_per_bin_sparse_binary( max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), include_diagonal=True, + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size = len(fingerprints) num_bins = len(selection_bins) @@ -216,6 +254,7 @@ def _compute_tanimoto_similarity_per_bin_sparse_binary( max_pairs_per_bin, selection_bins, include_diagonal, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -228,6 +267,7 @@ def _compute_tanimoto_similarity_per_bin_sparse_count( max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), include_diagonal=True, + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size = len(fingerprints_bins) num_bins = len(selection_bins) @@ -250,6 +290,7 @@ def _compute_tanimoto_similarity_per_bin_sparse_count( max_pairs_per_bin, selection_bins, include_diagonal, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -267,8 +308,10 @@ def compute_tanimoto_similarity_per_bin( fingerprint_type: str, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), include_diagonal=True, + random_seed=None, ) -> Tuple[np.ndarray, np.ndarray]: """Dispatch to the appropriate pairwise-per-bin Tanimoto implementation.""" + jit_random_seed = -1 if random_seed is None else int(random_seed) if fingerprint_type not in SUPPORTED_FINGERPRINT_TYPES: raise ValueError(f"Unsupported fingerprint type: {fingerprint_type}") @@ -278,6 +321,7 @@ def compute_tanimoto_similarity_per_bin( max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, include_diagonal=include_diagonal, + random_seed=jit_random_seed, ) if is_unfolded_binary_fingerprint_type(fingerprint_type): @@ -286,6 +330,7 @@ def compute_tanimoto_similarity_per_bin( max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, include_diagonal=include_diagonal, + random_seed=jit_random_seed, ) if is_unfolded_count_fingerprint_type(fingerprint_type): @@ -296,6 +341,7 @@ def compute_tanimoto_similarity_per_bin( max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, include_diagonal=include_diagonal, + random_seed=jit_random_seed, ) raise ValueError(f"Unsupported fingerprint type: {fingerprint_type}") @@ -310,21 +356,37 @@ def _fill_pairs_for_row_between_sets( target_offset, max_pairs_per_bin, selection_bins, + random_seed, ): + """Reservoir-sample bounded candidates for one cross-set row.""" num_bins = len(selection_bins) for bin_number in range(num_bins): selection_bin = selection_bins[bin_number] - indices = np.nonzero( - (tanimoto_scores > selection_bin[0]) & (tanimoto_scores <= selection_bin[1]) - )[0] + if random_seed >= 0: + np.random.seed( + (random_seed + 1_000_003 * row_index + 97 * bin_number) + % 4_294_967_295 + ) + + seen = 0 + filled = 0 + for candidate_index in range(len(tanimoto_scores)): + score = tanimoto_scores[candidate_index] + if not (score > selection_bin[0] and score <= selection_bin[1]): + continue - np.random.shuffle(indices) - indices = indices[:max_pairs_per_bin] - num_indices = len(indices) + seen += 1 + if filled < max_pairs_per_bin: + slot = filled + filled += 1 + else: + slot = np.random.randint(0, seen) + if slot >= max_pairs_per_bin: + continue - selected_pairs_per_bin[bin_number, row_index, :num_indices] = indices + target_offset - selected_scores_per_bin[bin_number, row_index, :num_indices] = tanimoto_scores[indices] + selected_pairs_per_bin[bin_number, row_index, slot] = candidate_index + target_offset + selected_scores_per_bin[bin_number, row_index, slot] = score @jit(nopython=True, parallel=True) @@ -333,6 +395,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_dense( fingerprints_2, max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size_1 = fingerprints_1.shape[0] size_2 = fingerprints_2.shape[0] @@ -353,6 +416,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_dense( size_1, max_pairs_per_bin, selection_bins, + random_seed, ) for idx_fingerprint_j in prange(size_2): @@ -368,6 +432,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_dense( 0, max_pairs_per_bin, selection_bins, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -379,6 +444,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_binary( fingerprints_2, max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size_1 = len(fingerprints_1) size_2 = len(fingerprints_2) @@ -399,6 +465,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_binary( size_1, max_pairs_per_bin, selection_bins, + random_seed, ) for idx_fingerprint_j in prange(size_2): @@ -414,6 +481,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_binary( 0, max_pairs_per_bin, selection_bins, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -427,6 +495,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_count( fingerprints_2_counts, max_pairs_per_bin, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), + random_seed=-1, ) -> Tuple[np.ndarray, np.ndarray]: size_1 = len(fingerprints_1_bins) size_2 = len(fingerprints_2_bins) @@ -453,6 +522,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_count( size_1, max_pairs_per_bin, selection_bins, + random_seed, ) for idx_fingerprint_j in prange(size_2): @@ -474,6 +544,7 @@ def _compute_tanimoto_similarity_per_bin_between_sets_sparse_count( 0, max_pairs_per_bin, selection_bins, + random_seed, ) return selected_pairs_per_bin, selected_scores_per_bin @@ -485,8 +556,10 @@ def compute_tanimoto_similarity_per_bin_between_sets( max_pairs_per_bin, fingerprint_type: str, selection_bins=np.array([(x / 10, x / 10 + 0.1) for x in range(10)], dtype=np.float32), + random_seed=None, ) -> Tuple[np.ndarray, np.ndarray]: """Compute cross-set Tanimoto per bin for all supported fingerprint types.""" + jit_random_seed = -1 if random_seed is None else int(random_seed) if fingerprint_type not in SUPPORTED_FINGERPRINT_TYPES: raise ValueError(f"Unsupported fingerprint type: {fingerprint_type}") @@ -496,6 +569,7 @@ def compute_tanimoto_similarity_per_bin_between_sets( fingerprints_2, max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, + random_seed=jit_random_seed, ) if is_unfolded_binary_fingerprint_type(fingerprint_type): @@ -504,6 +578,7 @@ def compute_tanimoto_similarity_per_bin_between_sets( fingerprints_2, max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, + random_seed=jit_random_seed, ) if is_unfolded_count_fingerprint_type(fingerprint_type): @@ -516,6 +591,7 @@ def compute_tanimoto_similarity_per_bin_between_sets( fingerprints_2_counts, max_pairs_per_bin=max_pairs_per_bin, selection_bins=selection_bins, + random_seed=jit_random_seed, ) raise ValueError(f"Unsupported fingerprint type: {fingerprint_type}") diff --git a/ms2deepscore/models/EmbeddingEvaluatorModel.py b/ms2deepscore/models/EmbeddingEvaluatorModel.py index ceb99332..9f6c5f63 100644 --- a/ms2deepscore/models/EmbeddingEvaluatorModel.py +++ b/ms2deepscore/models/EmbeddingEvaluatorModel.py @@ -152,15 +152,26 @@ def train_evaluator( with no_grad(): self.eval() val_losses = [] + for sample in val_generator: tanimoto_scores, ms2ds_scores, embeddings = sample + outputs = self(embeddings.reshape(-1, 1, embeddings.shape[-1]).to(device)) mse_per_embedding = ((tanimoto_scores - ms2ds_scores) ** 2).mean(dim=1) mse_per_embedding = mse_per_embedding.reshape(-1, 1).clone().detach() - loss = criterion(outputs.to(device), mse_per_embedding.to(device, dtype=float32)) - val_losses.append(loss_value) + loss = criterion( + outputs, + mse_per_embedding.to( + device, + dtype=float32, + ), + ) + + val_loss_value = loss.detach().item() + val_losses.append(val_loss_value) + print(f">>> Val_loss: {np.mean(val_losses):.6f}") self.train() diff --git a/ms2deepscore/pair_selection_cache.py b/ms2deepscore/pair_selection_cache.py new file mode 100644 index 00000000..a5457924 --- /dev/null +++ b/ms2deepscore/pair_selection_cache.py @@ -0,0 +1,231 @@ +"""Persistent caches for expensive training-pair preparation. + +Pair selection naturally separates into two reusable artifacts: + +1. Candidate pairs per Tanimoto bin (expensive all-vs-all fingerprint work). +2. The final balanced pair schedule (depends on balancing settings as well). + +Both artifacts are keyed by the structural metadata of the input spectra and by +only the settings that can affect the corresponding artifact. Arrays are stored +as .npy files so large candidate pools can be memory-mapped on reload. +""" + +from __future__ import annotations + +import hashlib +import json +import os +from pathlib import Path +from typing import Sequence + +import numpy as np + + +CACHE_SCHEMA_VERSION = 1 +CANDIDATE_ALGORITHM_VERSION = 1 +SELECTION_ALGORITHM_VERSION = 1 + + +def _json_hash(payload: dict) -> str: + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + return hashlib.sha256(encoded).hexdigest()[:24] + + +def _serialize_bins(score_bins) -> list[list[float]]: + return [[float(low), float(high)] for low, high in np.asarray(score_bins)] + + +def _spectrum_metadata_signature(spectra_sets: Sequence[Sequence]) -> str: + """Hash only metadata that can affect structure-based pair selection. + + Peak arrays deliberately are not hashed because they do not participate in + fingerprint/Tanimoto pair selection. Input order is included because legacy + representative-structure selection can use the first matching spectrum. + """ + digest = hashlib.sha256() + for set_index, spectra in enumerate(spectra_sets): + digest.update(f"SET:{set_index}\n".encode("ascii")) + for spectrum in spectra: + if spectrum is None: + payload = [None, None, None, None] + else: + payload = [ + spectrum.get("inchikey"), + spectrum.get("smiles"), + spectrum.get("inchi"), + spectrum.get("ionmode"), + ] + digest.update( + json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + ) + digest.update(b"\n") + return digest.hexdigest() + + +def _candidate_parameters(settings, mode: str, dataset_signature: str) -> dict: + return { + "schema_version": CACHE_SCHEMA_VERSION, + "algorithm_version": CANDIDATE_ALGORITHM_VERSION, + "mode": mode, + "dataset_signature": dataset_signature, + "fingerprint_type": settings.fingerprint_type, + "fingerprint_nbits": int(settings.fingerprint_nbits), + "max_pairs_per_bin": None + if settings.max_pairs_per_bin is None + else int(settings.max_pairs_per_bin), + "same_prob_bins": _serialize_bins(settings.same_prob_bins), + "include_diagonal": bool(settings.include_diagonal), + # Candidate subsampling uses random shuffling when a bin contains more + # than max_pairs_per_bin entries, so seed is part of artifact identity. + "random_seed": settings.random_seed, + } + + +def _selection_parameters(settings) -> dict: + return { + "schema_version": CACHE_SCHEMA_VERSION, + "algorithm_version": SELECTION_ALGORITHM_VERSION, + "average_inchikey_sampling_count": float(settings.average_inchikey_sampling_count), + "max_inchikey_sampling": int(settings.max_inchikey_sampling), + "max_pair_resampling": int(settings.max_pair_resampling), + } + + +class PairSelectionCache: + """Read/write persistent pair-selection artifacts under one cache root.""" + + def __init__(self, root: str | os.PathLike): + self.root = Path(root) + self.root.mkdir(parents=True, exist_ok=True) + + def _candidate_dir(self, spectra_sets, settings, mode: str): + dataset_signature = _spectrum_metadata_signature(spectra_sets) + params = _candidate_parameters(settings, mode, dataset_signature) + return self.root / f"candidates_{_json_hash(params)}", params + + @staticmethod + def _metadata_matches(path: Path, expected: dict) -> bool: + try: + with path.open("r", encoding="utf-8") as handle: + return json.load(handle) == expected + except (OSError, json.JSONDecodeError): + return False + + def load_candidates(self, spectra_sets, settings, mode: str): + candidate_dir, params = self._candidate_dir(spectra_sets, settings, mode) + metadata_path = candidate_dir / "metadata.json" + required = ( + candidate_dir / "inchikeys.npy", + candidate_dir / "available_pairs.npy", + candidate_dir / "available_scores.npy", + ) + if not self._metadata_matches(metadata_path, params) or not all(path.exists() for path in required): + return None + + inchikeys = np.load(required[0], allow_pickle=False).astype("U14").tolist() + available_pairs = np.load(required[1], mmap_mode="r", allow_pickle=False) + available_scores = np.load(required[2], mmap_mode="r", allow_pickle=False) + return inchikeys, available_pairs, available_scores, candidate_dir + + def save_candidates( + self, + spectra_sets, + settings, + mode: str, + inchikeys, + available_pairs: np.ndarray, + available_scores: np.ndarray, + ) -> Path: + candidate_dir, params = self._candidate_dir(spectra_sets, settings, mode) + candidate_dir.mkdir(parents=True, exist_ok=True) + + # InChIKey14 is ASCII. Store fixed-width bytes instead of NumPy Unicode + # (14 vs. 56 bytes per key) to keep large cached pair schedules compact. + np.save(candidate_dir / "inchikeys.npy", np.asarray(inchikeys, dtype="S14"), allow_pickle=False) + np.save( + candidate_dir / "available_pairs.npy", + np.asarray(available_pairs, dtype=np.int32), + allow_pickle=False, + ) + np.save( + candidate_dir / "available_scores.npy", + np.asarray(available_scores, dtype=np.float32), + allow_pickle=False, + ) + + tmp_metadata = candidate_dir / "metadata.json.tmp" + with tmp_metadata.open("w", encoding="utf-8") as handle: + json.dump(params, handle, indent=2, sort_keys=True) + tmp_metadata.replace(candidate_dir / "metadata.json") + return candidate_dir + + def _selection_dir(self, candidate_dir: Path, settings) -> tuple[Path, dict]: + params = _selection_parameters(settings) + return candidate_dir / f"selection_{_json_hash(params)}", params + + def load_selected_pairs(self, candidate_dir: Path, settings): + selection_dir, params = self._selection_dir(candidate_dir, settings) + metadata_path = selection_dir / "metadata.json" + pair1_path = selection_dir / "inchikey_1.npy" + pair2_path = selection_dir / "inchikey_2.npy" + scores_path = selection_dir / "scores.npy" + if not self._metadata_matches(metadata_path, params): + return None + if not (pair1_path.exists() and pair2_path.exists() and scores_path.exists()): + return None + + pair_1 = np.load(pair1_path, mmap_mode="r", allow_pickle=False) + pair_2 = np.load(pair2_path, mmap_mode="r", allow_pickle=False) + scores = np.load(scores_path, mmap_mode="r", allow_pickle=False) + if not (len(pair_1) == len(pair_2) == len(scores)): + return None + + # SpectrumPairGenerator's public API currently consumes tuples. Keeping + # that interface avoids a broad user-facing change while still skipping + # all expensive pair-selection work on cache hits. + return [ + ( + bytes(inchikey_1).decode("ascii"), + bytes(inchikey_2).decode("ascii"), + float(score), + ) + for inchikey_1, inchikey_2, score in zip(pair_1, pair_2, scores) + ] + + def save_selected_pairs(self, candidate_dir: Path, settings, selected_pairs) -> Path: + selection_dir, params = self._selection_dir(candidate_dir, settings) + selection_dir.mkdir(parents=True, exist_ok=True) + + n_pairs = len(selected_pairs) + inchikey_1 = np.empty(n_pairs, dtype="S14") + inchikey_2 = np.empty(n_pairs, dtype="S14") + scores = np.empty(n_pairs, dtype=np.float32) + for idx, (key_1, key_2, score) in enumerate(selected_pairs): + inchikey_1[idx] = key_1 + inchikey_2[idx] = key_2 + scores[idx] = score + + np.save(selection_dir / "inchikey_1.npy", inchikey_1, allow_pickle=False) + np.save(selection_dir / "inchikey_2.npy", inchikey_2, allow_pickle=False) + np.save(selection_dir / "scores.npy", scores, allow_pickle=False) + + tmp_metadata = selection_dir / "metadata.json.tmp" + with tmp_metadata.open("w", encoding="utf-8") as handle: + json.dump(params, handle, indent=2, sort_keys=True) + tmp_metadata.replace(selection_dir / "metadata.json") + return selection_dir + + +def resolve_pair_selection_cache_directory(settings, fallback_root=None): + """Return the shared cache directory requested by training settings.""" + if not getattr(settings, "use_pair_selection_cache", True): + return None + explicit = getattr(settings, "pair_selection_cache_directory", None) + if explicit is not None: + return explicit + results_folder = getattr(settings, "results_folder", None) + if results_folder is not None: + return os.path.join(results_folder, "pair_selection_cache") + if fallback_root is not None: + return os.path.join(fallback_root, "pair_selection_cache") + return None diff --git a/ms2deepscore/train_new_model/DataGeneratorEmbeddingEvaluation.py b/ms2deepscore/train_new_model/DataGeneratorEmbeddingEvaluation.py index 19ccf4a0..5a46d229 100644 --- a/ms2deepscore/train_new_model/DataGeneratorEmbeddingEvaluation.py +++ b/ms2deepscore/train_new_model/DataGeneratorEmbeddingEvaluation.py @@ -1,20 +1,19 @@ from typing import List import numpy as np -import pandas as pd -from torch import tensor +from torch import no_grad, tensor from matchms import Spectrum -from matchms.similarity.vector_similarity_functions import jaccard_similarity_matrix from ms2deepscore.SettingsMS2Deepscore import SettingsEmbeddingEvaluator from ms2deepscore.models import SiameseSpectralModel from ms2deepscore.tensorize_spectra import tensorize_spectra from ms2deepscore.train_new_model.inchikey_pair_selection import compute_fingerprints_for_training +from ms2deepscore.fingerprint_similarity_computations import compute_fingerprint_similarity_matrix from ms2deepscore.vector_operations import cosine_similarity_matrix class DataGeneratorEmbeddingEvaluation: - """Generates data for training an embedding evaluation model. + """Generate data for training an embedding-evaluation model. This class provides a data for the training of an embedding evaluation model. It follows a simple strategy: iterate through all spectra and randomly pick another @@ -37,7 +36,6 @@ def __init__( device="cpu", ): """ - Parameters ---------- spectrums @@ -48,21 +46,25 @@ def __init__( self.current_index = 0 self.settings = settings self.spectrums = spectrums - self.inchikey14s = [s.get("inchikey")[:14] for s in spectrums] # type: ignore + self.inchikey14s = [s.get("inchikey")[:14] for s in spectrums] self.ms2ds_model = ms2ds_model self.device = device self.ms2ds_model.to(self.device) + self.ms2ds_model.eval() self.indexes = np.arange(len(self.spectrums)) self.batch_size = self.settings.evaluator_distribution_size - self.fingerprint_df = self.compute_fingerprint_dataframe( + + self.fingerprint_type = self.ms2ds_model.model_settings.fingerprint_type + 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, + self.fingerprint_type, + self.ms2ds_model.model_settings.fingerprint_nbits, ) + self.fingerprint_index_by_inchikey = { + inchikey: idx for idx, inchikey in enumerate(fingerprint_inchikeys) + } - # Initialize random number generator self.rng = np.random.default_rng(self.settings.random_seed) - self.on_epoch_end() def __len__(self): @@ -76,10 +78,22 @@ def __next__(self): batch = self.__getitem__(self.current_index) self.current_index += 1 return batch - self.current_index = 0 # make generator executable again + self.current_index = 0 self.on_epoch_end() raise StopIteration + def _select_fingerprints(self, inchikeys): + try: + positions = [self.fingerprint_index_by_inchikey[key] for key in inchikeys] + except KeyError as exc: + raise ValueError( + f"No fingerprint available for InChIKey {exc.args[0]!r}." + ) from exc + + if isinstance(self.fingerprints, np.ndarray): + return self.fingerprints[positions] + return [self.fingerprints[position] for position in positions] + def _compute_embeddings_and_scores(self, batch_index: int): batch_size = self.batch_size indexes = self.indexes[batch_index * batch_size : ((batch_index + 1) * batch_size)] @@ -87,43 +101,27 @@ 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 no_grad(): + embeddings = self.ms2ds_model.encoder( + spec_tensors.to(self.device), meta_tensors.to(self.device) + ) + embeddings_cpu = embeddings.detach().cpu() - ms2ds_scores = cosine_similarity_matrix(embeddings.cpu().detach().numpy(), embeddings.cpu().detach().numpy()) + embedding_array = embeddings_cpu.numpy() + ms2ds_scores = cosine_similarity_matrix(embedding_array, embedding_array) - # Compute true scores inchikeys = [self.inchikey14s[i] for i in indexes] - fingerprints = self.fingerprint_df.loc[inchikeys].to_numpy() - - tanimoto_scores = jaccard_similarity_matrix(fingerprints, fingerprints) + fingerprints = self._select_fingerprints(inchikeys) + tanimoto_scores = compute_fingerprint_similarity_matrix( + fingerprints, + fingerprints, + fingerprint_type=self.fingerprint_type, + ) - return tensor(tanimoto_scores), tensor(ms2ds_scores), embeddings.cpu().detach() + return tensor(tanimoto_scores), tensor(ms2ds_scores), embeddings_cpu def on_epoch_end(self): - """Updates indexes after each epoch.""" self.rng.shuffle(self.indexes) def __getitem__(self, batch_index: int): - """Generate one batch of data.""" return self._compute_embeddings_and_scores(batch_index) - - def compute_fingerprint_dataframe( - self, - spectrums: List[Spectrum], - fingerprint_type, - fingerprint_nbits, - ) -> pd.DataFrame: - """Returns a dataframe with a fingerprints dataframe - - spectrums: - A list of spectra - settings: - The settings that should be used for selecting the compound pairs wrapper. The settings should be specified as a - SettingsMS2Deepscore object. - """ - fingerprints, inchikeys14_unique = compute_fingerprints_for_training( - spectrums, fingerprint_type, fingerprint_nbits - ) - - fingerprints_df = pd.DataFrame(fingerprints, index=inchikeys14_unique) - return fingerprints_df diff --git a/ms2deepscore/train_new_model/TrainingBatchGenerator.py b/ms2deepscore/train_new_model/TrainingBatchGenerator.py index 71fb1f27..51b488a1 100644 --- a/ms2deepscore/train_new_model/TrainingBatchGenerator.py +++ b/ms2deepscore/train_new_model/TrainingBatchGenerator.py @@ -92,8 +92,8 @@ def __getitem__(self, batch_index: int): """ if self.model_settings.use_fixed_set and batch_index in self.fixed_set: return self.fixed_set[batch_index] - if self.model_settings.random_seed is not None and batch_index == 0: - self.rng = np.random.default_rng(self.model_settings.random_seed) + # Seed once in __init__. Re-seeding at batch 0 made every epoch use + # exactly the same augmentation sequence whenever random_seed was set. spectrum_pairs = self._spectrum_pair_generator() spectra_1, spectra_2, meta_1, meta_2, targets = self._tensorize_all(spectrum_pairs) diff --git a/ms2deepscore/train_new_model/inchikey_pair_selection.py b/ms2deepscore/train_new_model/inchikey_pair_selection.py index d4625635..b237a3e0 100644 --- a/ms2deepscore/train_new_model/inchikey_pair_selection.py +++ b/ms2deepscore/train_new_model/inchikey_pair_selection.py @@ -6,55 +6,108 @@ from tqdm import tqdm from ms2deepscore.SettingsMS2Deepscore import SettingsMS2Deepscore from ms2deepscore.train_new_model import SpectrumPairGenerator -from ms2deepscore.fingerprint_utils import derive_fingerprint_from_smiles_or_inchi +from ms2deepscore.fingerprint_utils import ( + derive_fingerprint_from_smiles, + normalize_to_smiles, +) +from ms2deepscore.pair_selection_cache import ( + PairSelectionCache, + resolve_pair_selection_cache_directory, +) from ms2deepscore.fingerprint_similarity_computations import compute_tanimoto_similarity_per_bin def create_spectrum_pair_generator( spectra: List[Spectrum], settings: SettingsMS2Deepscore, + cache_directory=None, ) -> SpectrumPairGenerator: - """Returns a SpectrumPairGenerator object containing equally balanced pairs over the different bins + """Return a balanced SpectrumPairGenerator, optionally using persistent caches. - spectra: - A list of spectra - settings: - The settings that should be used for selecting the compound pairs wrapper. The settings should be specified as a - SettingsMS2Deepscore object. - - Returns - ------- - SpectrumPairGenerator - SpectrumPairGenerator containing balanced pairs. The pairs are stored as [(inchikey1, inchikey2, score)] + The persistent cache stores both the expensive per-bin candidate pool and + the final balanced pair schedule. Cache identity contains only structural + spectrum metadata plus the settings that can influence each artifact. """ - if settings.random_seed is not None: - np.random.seed(settings.random_seed) + if cache_directory is None: + cache_directory = resolve_pair_selection_cache_directory(settings) + cache = PairSelectionCache(cache_directory) if cache_directory is not None else None + candidate_dir = None + + if cache is not None: + cached = cache.load_candidates((spectra,), settings, mode="same_set") + else: + cached = None + + if cached is not None: + inchikeys14_unique, available_pairs_per_bin_matrix, available_scores_per_bin_matrix, candidate_dir = cached + selected_pairs = cache.load_selected_pairs(candidate_dir, settings) + if selected_pairs is not None: + print(f"Reusing {len(selected_pairs)} cached training compound pairs from {candidate_dir}") + return SpectrumPairGenerator( + selected_pairs, spectra, settings.shuffle, settings.random_seed + ) + print(f"Reusing cached Tanimoto candidate pairs from {candidate_dir}") + else: + fingerprints, inchikeys14_unique = compute_fingerprints_for_training( + spectra, + settings.fingerprint_type, + settings.fingerprint_nbits, + ) - fingerprints, inchikeys14_unique = compute_fingerprints_for_training( - spectra, - settings.fingerprint_type, - settings.fingerprint_nbits + if len(inchikeys14_unique) < settings.batch_size: + raise ValueError("The number of unique inchikeys must be larger than the batch size.") + + max_pairs_per_bin = settings.max_pairs_per_bin + if max_pairs_per_bin is None: + # Honor the documented setting. This can be extremely memory hungry, + # so bounded max_pairs_per_bin remains strongly recommended. + max_pairs_per_bin = len(inchikeys14_unique) + + available_pairs_per_bin_matrix, available_scores_per_bin_matrix = compute_tanimoto_similarity_per_bin( + fingerprints, + max_pairs_per_bin, + fingerprint_type=settings.fingerprint_type, + selection_bins=settings.same_prob_bins, + include_diagonal=settings.include_diagonal, + random_seed=settings.random_seed, ) + if cache is not None: + candidate_dir = cache.save_candidates( + (spectra,), + settings, + mode="same_set", + inchikeys=inchikeys14_unique, + available_pairs=available_pairs_per_bin_matrix, + available_scores=available_scores_per_bin_matrix, + ) + # Re-open as memory maps so the rest of the workflow does not need + # another full in-memory copy of the cached arrays. + cached = cache.load_candidates((spectra,), settings, mode="same_set") + if cached is not None: + inchikeys14_unique, available_pairs_per_bin_matrix, available_scores_per_bin_matrix, candidate_dir = cached + if len(inchikeys14_unique) < settings.batch_size: raise ValueError("The number of unique inchikeys must be larger than the batch size.") - available_pairs_per_bin_matrix, available_scores_per_bin_matrix = compute_tanimoto_similarity_per_bin( - fingerprints, - settings.max_pairs_per_bin, - fingerprint_type=settings.fingerprint_type, - selection_bins=settings.same_prob_bins, - include_diagonal=settings.include_diagonal, - ) pair_frequency_matrixes = balanced_selection_of_pairs_per_bin( - available_pairs_per_bin_matrix, settings) + available_pairs_per_bin_matrix, settings + ) selected_pairs_per_bin = convert_to_selected_pairs_list( - pair_frequency_matrixes, available_pairs_per_bin_matrix, - available_scores_per_bin_matrix, inchikeys14_unique) + pair_frequency_matrixes, + available_pairs_per_bin_matrix, + available_scores_per_bin_matrix, + inchikeys14_unique, + ) + selected_pairs = [pair for pairs in selected_pairs_per_bin for pair in pairs] - return SpectrumPairGenerator([pair for pairs in selected_pairs_per_bin for pair in pairs], - spectra, settings.shuffle, settings.random_seed) + if cache is not None and candidate_dir is not None: + cache.save_selected_pairs(candidate_dir, settings, selected_pairs) + + return SpectrumPairGenerator( + selected_pairs, spectra, settings.shuffle, settings.random_seed + ) def compute_fingerprints_for_training( @@ -92,21 +145,25 @@ def compute_fingerprints_for_training( structure = spectrum.get("smiles") if structure is None: structure = spectrum.get("inchi") - if structure is None: continue - structure_list.append(structure) + # Normalize InChI before appending the matching inchikey. + normalized_structure = normalize_to_smiles(structure) + if normalized_structure is None: + continue + + structure_list.append(normalized_structure) valid_inchikeys.append(inchikey14) if len(structure_list) == 0: raise ValueError("No valid SMILES/InChI entries available for fingerprint calculation") - fingerprints = derive_fingerprint_from_smiles_or_inchi( + fingerprints = derive_fingerprint_from_smiles( structure_list, fingerprint_type=fingerprint_type, nbits=nbits, - policy_invalid="keep", + policy_invalid_smiles="keep", ) if len(fingerprints) == 0: @@ -190,7 +247,14 @@ def convert_to_selected_pairs_list(pair_frequency_matrixes: np.ndarray, available_pairs_per_bin_matrix: np.ndarray, scores_matrix: np.ndarray, inchikeys14_unique: List[str]): - """Convert the matrixes denoting the pairs to a list of pairs, encoded as [(inchikey1, inchikey2, score)] + """Convert pair frequencies to ``(inchikey1, inchikey2, score)`` lists. + + The previous implementation (version<=0.29) iterated in Python over every slot of the + dense ``(bins, compounds, max_pairs_per_bin)`` candidate cube. At large + scale that can mean hundreds of millions of Python-loop iterations even + though only a small fraction of slots have a non-zero selected frequency. + Here NumPy finds non-zero entries in C and canonical pair IDs remove the + mirrored duplicates before the much smaller Python expansion step. Parameters ---------- @@ -211,26 +275,53 @@ def convert_to_selected_pairs_list(pair_frequency_matrixes: np.ndarray, This is used to map the indexes of inchikeys used in the matrixes, to the corresponding inchikeys. """ selected_pairs_per_bin = [] - for bin_id, bin_pair_frequency_matrix in enumerate(tqdm(pair_frequency_matrixes)): + nr_of_inchikeys = len(inchikeys14_unique) + + for bin_id in tqdm(range(pair_frequency_matrixes.shape[0])): + frequency_matrix = pair_frequency_matrixes[bin_id] + row_indices, column_indices = np.nonzero(frequency_matrix > 0) + + if len(row_indices) == 0: + selected_pairs_per_bin.append([]) + continue + + partner_indices = available_pairs_per_bin_matrix[ + bin_id, row_indices, column_indices + ] + valid = partner_indices >= 0 + row_indices = row_indices[valid] + column_indices = column_indices[valid] + partner_indices = partner_indices[valid] + + lower = np.minimum(row_indices, partner_indices).astype(np.int64) + upper = np.maximum(row_indices, partner_indices).astype(np.int64) + pair_codes = lower * nr_of_inchikeys + upper + + # ``np.unique`` sorts by code; sort the returned first-occurrence + # positions to retain the legacy row-major ordering as closely as + # possible while dropping mirrored duplicates. + _, first_occurrences = np.unique(pair_codes, return_index=True) + first_occurrences.sort() + selected_pairs = [] - for inchikey1_index, pair_frequency_row in enumerate(bin_pair_frequency_matrix): - for column_index, pair_frequency in enumerate(pair_frequency_row): - if pair_frequency > 0: - inchikey2_index = available_pairs_per_bin_matrix[bin_id][inchikey1_index][column_index] - score = scores_matrix[bin_id][inchikey1_index][column_index] - # This ensures that the order is the same. - # This is important for the cross ionization mode selection. - if inchikey1_index < inchikey2_index: - selected_pairs.extend( - [(inchikeys14_unique[inchikey1_index], inchikeys14_unique[inchikey2_index], score)] * pair_frequency) - else: - selected_pairs.extend( - [(inchikeys14_unique[inchikey2_index], inchikeys14_unique[inchikey1_index], score)] * pair_frequency) - # remove duplicate pairs - position_of_first_inchikey_in_matrix = available_pairs_per_bin_matrix[bin_id][ - inchikey2_index] == inchikey1_index - bin_pair_frequency_matrix[inchikey2_index][position_of_first_inchikey_in_matrix] = 0 + for position in first_occurrences: + inchikey1_index = int(lower[position]) + inchikey2_index = int(upper[position]) + column_index = int(column_indices[position]) + original_row = int(row_indices[position]) + pair_frequency = int(frequency_matrix[original_row, column_index]) + score = float(scores_matrix[bin_id, original_row, column_index]) + + selected_pairs.extend( + [( + inchikeys14_unique[inchikey1_index], + inchikeys14_unique[inchikey2_index], + score, + )] * pair_frequency + ) + selected_pairs_per_bin.append(selected_pairs) + return selected_pairs_per_bin @@ -269,12 +360,20 @@ def select_balanced_pairs(available_pairs_for_bin_matrix: np.ndarray, """ num_inchikeys = available_pairs_for_bin_matrix.shape[0] - # Initialize pair frequency matrix - pair_frequency = np.zeros_like(available_pairs_for_bin_matrix, dtype=int) + # Initialize pair frequencies with the smallest safe integer dtype. + sentinel_value = 2 * max_resampling + frequency_dtype = ( + np.int32 + if sentinel_value <= np.iinfo(np.int32).max + else np.int64 + ) + pair_frequency = np.zeros_like( + available_pairs_for_bin_matrix, dtype=frequency_dtype + ) # Mask for invalid pairs (where value is -1) invalid_mask = (available_pairs_for_bin_matrix == -1) - pair_frequency[invalid_mask] = max_resampling * 2 # Ensure these pairs are never selected + pair_frequency[invalid_mask] = sentinel_value # Ensure these pairs are never selected # Initialize available inchikeys as a min-heap based on inchikey_counts available_inchikey_indexes = [(inchikey_counts[i], i) for i in range(num_inchikeys) @@ -320,11 +419,15 @@ def select_balanced_pairs(available_pairs_for_bin_matrix: np.ndarray, if not np.any(valid_pairs_mask & valid_inchikeys_mask): continue # No valid pairs left for this inchikey - # Among valid pairs, find those with the lowest pair frequency - min_pair_freq = np.min(pair_freq_row[valid_pairs_mask & valid_inchikeys_mask]) - min_freq_mask = pair_freq_row == min_pair_freq + # Among valid pairs, find those with the lowest pair frequency. + # Keep the validity mask when resolving ties: the previous code + # could re-introduce partners that were already above + # max_inchikey_count merely because they had the same pair count. + candidate_mask = valid_pairs_mask & valid_inchikeys_mask + min_pair_freq = np.min(pair_freq_row[candidate_mask]) + min_freq_mask = candidate_mask & (pair_freq_row == min_pair_freq) - # From the least resampled inchikey select the leas sampled inchikey + # From the least-resampled pairs select the least-sampled partner. min_inchikey_count_idx = np.argmin(second_inchikey_counts[min_freq_mask]) second_inchikey_with_lowest_count = available_pairs_row[min_freq_mask][min_inchikey_count_idx] @@ -395,6 +498,6 @@ def select_inchi_for_unique_inchikeys( # ID of the spectrum with the most frequent inchi ID = idx[np.where(inchi_array[idx] == most_common_inchi)[0][0]] - spectra_selected.append(list_of_spectra[ID].clone()) + spectra_selected.append(list_of_spectra[ID]) return spectra_selected, inchikeys14_unique diff --git a/ms2deepscore/train_new_model/train_ms2deepscore.py b/ms2deepscore/train_new_model/train_ms2deepscore.py index 19ddca10..c65b75dc 100644 --- a/ms2deepscore/train_new_model/train_ms2deepscore.py +++ b/ms2deepscore/train_new_model/train_ms2deepscore.py @@ -15,6 +15,8 @@ from ms2deepscore.train_new_model.inchikey_pair_selection_cross_ionmode import create_data_generator_across_ionmodes from ms2deepscore.validation_loss_calculation.ValidationLossCalculator import \ ValidationLossCalculator +from ms2deepscore.pair_selection_cache import resolve_pair_selection_cache_directory +from ms2deepscore.models.load_model import load_model def train_ms2ds_model( @@ -31,7 +33,11 @@ def train_ms2ds_model( if settings.balanced_sampling_across_ionmodes: train_generator = create_data_generator_across_ionmodes(training_spectra, settings=settings) else: - spectrum_pair_generator = create_spectrum_pair_generator(training_spectra, settings=settings) + spectrum_pair_generator = create_spectrum_pair_generator( + training_spectra, + settings=settings, + cache_directory=resolve_pair_selection_cache_directory(settings, fallback_root=results_folder), + ) train_generator = TrainingBatchGenerator(spectrum_pair_generator=spectrum_pair_generator, settings=settings) # Create a validation loss calculator validation_loss_calculator = ValidationLossCalculator(validation_spectra, @@ -49,6 +55,12 @@ def train_ms2ds_model( patience=settings.patience, loss_function=settings.loss_function, checkpoint_filename=output_model_file_name, lambda_l1=0, lambda_l2=0) + # ``train`` checkpoints the best validation model, while ``model`` contains + # the weights from the final epoch. Downstream consumers (notably the + # embedding evaluator) should use the same best weights that were persisted. + if os.path.isfile(output_model_file_name): + model = load_model(output_model_file_name) + model.export_to_onnx(results_folder) return model, history diff --git a/ms2deepscore/validation_loss_calculation/ValidationLossCalculator.py b/ms2deepscore/validation_loss_calculation/ValidationLossCalculator.py index 9daa0e05..a33f0549 100644 --- a/ms2deepscore/validation_loss_calculation/ValidationLossCalculator.py +++ b/ms2deepscore/validation_loss_calculation/ValidationLossCalculator.py @@ -19,6 +19,7 @@ def __init__( val_spectrums, settings: SettingsMS2Deepscore, chunk_size: int = 10_000, + tanimoto_scores: pd.DataFrame | None = None, ): """ Parameters: @@ -38,12 +39,15 @@ def __init__( self.inchikeys14 = [spectrum.get("inchikey")[:14] for spectrum in self.val_spectrums] self.unique_inchikeys14 = sorted(set(self.inchikeys14)) - self.tanimoto_scores = calculate_tanimoto_scores_unique_inchikey( - list_of_spectra_1=self.val_spectrums, - list_of_spectra_2=None, - fingerprint_type=self.settings.fingerprint_type, - nbits=self.settings.fingerprint_nbits - ) + if tanimoto_scores is None: + self.tanimoto_scores = calculate_tanimoto_scores_unique_inchikey( + list_of_spectra_1=self.val_spectrums, + list_of_spectra_2=None, + fingerprint_type=self.settings.fingerprint_type, + nbits=self.settings.fingerprint_nbits, + ) + else: + self.tanimoto_scores = tanimoto_scores @staticmethod def _chunk_indices(n_items: int, chunk_size: int): @@ -213,6 +217,18 @@ def _compute_global_average_losses_per_inchikey_pair(self, embeddings, loss_type block_sum, block_count, ) + # Only upper-triangular chunk pairs are computed. Mirror + # off-diagonal blocks so every directed InChIKey pair is + # represented consistently. Otherwise pairs that happen to + # fall in the same chunk were counted in both directions, + # while cross-chunk pairs were present in only one. + if chunk_idx_i != chunk_idx_j: + self._update_global_pair_accumulators( + loss_sum_per_type[loss_type], + loss_count_per_type[loss_type], + block_sum.T, + block_count.T, + ) average_loss_per_type = {} for loss_type in loss_types: diff --git a/ms2deepscore/wrapper_functions/training_wrapper_functions.py b/ms2deepscore/wrapper_functions/training_wrapper_functions.py index 59f718b9..dbabeac7 100644 --- a/ms2deepscore/wrapper_functions/training_wrapper_functions.py +++ b/ms2deepscore/wrapper_functions/training_wrapper_functions.py @@ -15,6 +15,9 @@ train) from ms2deepscore.SettingsMS2Deepscore import SettingsMS2Deepscore, SettingsEmbeddingEvaluator from ms2deepscore.train_new_model import TrainingBatchGenerator, create_spectrum_pair_generator +from ms2deepscore.train_new_model.inchikey_pair_selection_cross_ionmode import ( + create_data_generator_across_ionmodes, +) from ms2deepscore.validation_loss_calculation.ValidationLossCalculator import ValidationLossCalculator from ms2deepscore.train_new_model.train_ms2deepscore import \ train_ms2ds_model, plot_history, save_history @@ -23,6 +26,7 @@ from ms2deepscore.utils import load_spectra_as_list from ms2deepscore.wrapper_functions.plotting_wrapper_functions import \ create_plots_between_ionmodes +from ms2deepscore.pair_selection_cache import resolve_pair_selection_cache_directory def train_ms2deepscore_wrapper(settings: SettingsMS2Deepscore, @@ -92,8 +96,6 @@ def parameter_search( """ print("Initialize Stored Data") split_data_if_necessary(base_settings) - validation_spectra = load_spectra_in_ionmode(base_settings.validation_spectra_file_name, base_settings.ionisation_mode) - print("Load training data") # Split training in pos and neg and create val and training split and select for the right ionisation mode. training_spectra = load_spectra_in_ionmode(base_settings.training_spectra_file_name, base_settings.ionisation_mode) @@ -106,7 +108,7 @@ def parameter_search( negative_validation_spectra = load_spectra_in_ionmode(base_settings.validation_spectra_file_name, "negative") results = {} - train_generator = None + validation_tanimoto_cache = {} # Generate all combinations of setting variations keys, values = zip(*setting_variations.items()) @@ -115,34 +117,45 @@ def parameter_search( settings_dict = base_settings.get_dict() settings_dict.update(params) settings = SettingsMS2Deepscore(**settings_dict) - settings.time_stamp = datetime.now().strftime("%Y_%m_%d_%H_%M_%S") - + # base_settings.get_dict() contains the already-derived model directory. + # Re-derive it for every combination or parameter sweeps overwrite each + # other. Microseconds avoid collisions for fast successive runs. + settings.time_stamp = datetime.now().strftime("%Y_%m_%d_%H_%M_%S_%f") + settings.model_directory_name = os.path.join( + settings.results_folder, settings.create_model_directory_name() + ) print(f"Testing combination: {params}") - # TODO (mabye): implement smarter way to now always re-initialize the generators - # fields_affecting_generators = [ - # "fingerprint_type", - # "fingerprint_nbits", - # "max_pairs_per_bin", - # "same_prob_bins", - # "include_diagonal", - # "mz_bin_width", - # ] - # search_includes_generator_parameters = False - # for field in fields_affecting_generators: - # if field in keys: - # search_includes_generator_parameters = True - # if search_includes_generator_parameters or (train_generator is None): # Make folder and save settings os.makedirs(settings.model_directory_name, exist_ok=True) settings.save_to_file(os.path.join(settings.model_directory_name, "settings.json")) - # Create a training generator - spectrum_pair_generator = create_spectrum_pair_generator(training_spectra, settings=settings) - train_generator = TrainingBatchGenerator(spectrum_pair_generator=spectrum_pair_generator, settings=settings) - # Create a validation loss calculator - validation_loss_calculator = ValidationLossCalculator(validation_spectra, - settings=settings) + # Create a training generator. The expensive compound-pair preparation + # is content-keyed and persisted across parameter combinations/runs. + if settings.balanced_sampling_across_ionmodes: + train_generator = create_data_generator_across_ionmodes( + training_spectra, settings=settings + ) + else: + spectrum_pair_generator = create_spectrum_pair_generator( + training_spectra, + settings=settings, + cache_directory=resolve_pair_selection_cache_directory(settings), + ) + train_generator = TrainingBatchGenerator( + spectrum_pair_generator=spectrum_pair_generator, settings=settings + ) + + # Validation Tanimoto targets depend on fingerprint settings, not on the + # neural-network hyperparameters. Reuse them within the sweep. + validation_key = (settings.fingerprint_type, settings.fingerprint_nbits) + cached_tanimoto = validation_tanimoto_cache.get(validation_key) + validation_loss_calculator = ValidationLossCalculator( + validation_spectra, settings=settings, tanimoto_scores=cached_tanimoto + ) + validation_tanimoto_cache.setdefault( + validation_key, validation_loss_calculator.tanimoto_scores + ) model = SiameseSpectralModel(settings=settings) @@ -171,7 +184,10 @@ def parameter_search( scores_between_all_ionmodes = CalculateScoresBetweenAllIonmodes( model_file_name=os.path.join(settings.model_directory_name, settings.model_file_name), positive_validation_spectra=positive_validation_spectra, - negative_validation_spectra=negative_validation_spectra) + negative_validation_spectra=negative_validation_spectra, + fingerprint_type=settings.fingerprint_type, + n_bits_fingerprint=settings.fingerprint_nbits, + ) combination_results = { "params": params, diff --git a/tests/conftest.py b/tests/conftest.py index 68b9fd7c..957fda1a 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,2 +1,98 @@ import matplotlib -matplotlib.use("Agg", force=True) \ No newline at end of file +import pytest +import torch +from torch import nn + +from ms2deepscore.SettingsMS2Deepscore import ( + SettingsEmbeddingEvaluator, + SettingsMS2Deepscore, +) +from ms2deepscore.train_new_model.DataGeneratorEmbeddingEvaluation import ( + DataGeneratorEmbeddingEvaluation, +) +from tests.create_test_spectra import create_test_spectra +matplotlib.use("Agg", force=True) + + +class DeterministicMockEncoder(nn.Module): + """Small deterministic encoder implementing the interface used by the generator. + + Besides returning stable embeddings, it records whether inference happened in + eval mode and with gradients disabled. This makes the mock useful for testing + inference semantics without constructing a full SiameseSpectralModel. + """ + + def __init__(self, embedding_dim: int = 128): + super().__init__() + self.embedding_dim = embedding_dim + self.last_training = None + self.last_grad_enabled = None + + def forward(self, spec_tensors, meta_tensors): + self.last_training = self.training + self.last_grad_enabled = torch.is_grad_enabled() + + # Produce deterministic, non-zero embeddings for every input spectrum. + # The weighted spectral sum makes embeddings spectrum-dependent, while the + # offsets prevent zero vectors (important for cosine similarities). + weights = torch.linspace( + 0.5, + 1.5, + spec_tensors.shape[1], + dtype=spec_tensors.dtype, + device=spec_tensors.device, + ) + summary = (spec_tensors * weights).sum(dim=1, keepdim=True) + if meta_tensors.shape[1] > 0: + summary = summary + meta_tensors.sum(dim=1, keepdim=True) + + offsets = torch.linspace( + 0.01, + 1.0, + self.embedding_dim, + dtype=spec_tensors.dtype, + device=spec_tensors.device, + ).unsqueeze(0) + return summary + offsets + + +class MockMS2DSModel(nn.Module): + """Minimal, valid nn.Module satisfying DataGeneratorEmbeddingEvaluation.""" + + def __init__(self, embedding_dim: int = 128): + super().__init__() + self.model_settings = SettingsMS2Deepscore( + embedding_dim=embedding_dim, + fingerprint_nbits=128, + ) + self.encoder = DeterministicMockEncoder(embedding_dim=embedding_dim) + + +@pytest.fixture +def mock_ms2ds_model(): + return MockMS2DSModel() + + +@pytest.fixture +def embedding_evaluator_generator_settings(): + return SettingsEmbeddingEvaluator( + evaluator_distribution_size=10, + random_seed=123, + ) + + +@pytest.fixture +def data_generator_embedding_evaluation( + mock_ms2ds_model, + embedding_evaluator_generator_settings, +): + spectra = create_test_spectra( + num_of_unique_inchikeys=25, + num_of_spectra_per_inchikey=2, + ) + return DataGeneratorEmbeddingEvaluation( + spectrums=spectra, + ms2ds_model=mock_ms2ds_model, + settings=embedding_evaluator_generator_settings, + device="cpu", + ) diff --git a/tests/test_data_generators.py b/tests/test_data_generators.py index 6b84bb2e..119f3dc3 100644 --- a/tests/test_data_generators.py +++ b/tests/test_data_generators.py @@ -1,14 +1,20 @@ -import pytest -import numpy as np -from torch import rand, Size from collections import Counter + +import numpy as np +import pytest +import torch from matchms import Spectrum -from ms2deepscore.SettingsMS2Deepscore import SettingsMS2Deepscore, SettingsEmbeddingEvaluator -from ms2deepscore.models import SiameseSpectralModel + +from ms2deepscore.SettingsMS2Deepscore import SettingsEmbeddingEvaluator, SettingsMS2Deepscore from ms2deepscore.tensorize_spectra import tensorize_spectra +from ms2deepscore.train_new_model import ( + SpectrumPairGenerator, + create_spectrum_pair_generator, +) +from ms2deepscore.train_new_model.DataGeneratorEmbeddingEvaluation import ( + DataGeneratorEmbeddingEvaluation, +) from ms2deepscore.train_new_model.TrainingBatchGenerator import TrainingBatchGenerator -from ms2deepscore.train_new_model.DataGeneratorEmbeddingEvaluation import DataGeneratorEmbeddingEvaluation -from ms2deepscore.train_new_model import SpectrumPairGenerator, create_spectrum_pair_generator from ms2deepscore.train_new_model.inchikey_pair_selection_cross_ionmode import ( create_data_generator_across_ionmodes, select_compound_pairs_wrapper_across_ionmode, @@ -16,300 +22,333 @@ from tests.create_test_spectra import create_test_spectra -class MockMS2DSModel(SiameseSpectralModel): - def __init__(self): - self.model_settings = SettingsMS2Deepscore() - - def encoder(self, spec_tensors, meta_tensors): - # Return mock embeddings as random tensors - return rand(spec_tensors.size(0), 128) # Assuming embedding size of 128 - - def to(self, device): - pass - - -@pytest.fixture -def data_generator_embedding_evaluation(): - spectrums = create_test_spectra(num_of_unique_inchikeys=25, num_of_spectra_per_inchikey=2) - params = {"evaluator_distribution_size": 10} - return DataGeneratorEmbeddingEvaluation( - spectrums=spectrums, ms2ds_model=MockMS2DSModel(), settings=SettingsEmbeddingEvaluator(**params), device="cpu" - ) - +SELECTED_PAIRS = [ + ("CCCCCCCCCCCCCC", "DDDDDDDDDDDDDD", 0.25), + ("BBBBBBBBBBBBBB", "DDDDDDDDDDDDDD", 0.6666667), + ("AAAAAAAAAAAAAA", "CCCCCCCCCCCCCC", 1.0), + ("AAAAAAAAAAAAAA", "BBBBBBBBBBBBBB", 0.33333334), +] -def collect_results(generator, batch_size, dimension): - n_batches = len(generator) - X = np.zeros((batch_size, dimension, 2, n_batches)) - y = np.zeros((batch_size, n_batches)) - for i, batch in enumerate(generator): - X[:, :, 0, i] = batch[0][0] - X[:, :, 1, i] = batch[0][1] - y[:, i] = batch[1] - return X, y - -def test_tensorize_spectra(): - spectrum = Spectrum(mz=np.array([10, 500, 999.9]), intensities=np.array([0.5, 0.5, 1])) - settings = SettingsMS2Deepscore( - min_mz=10, max_mz=1000, mz_bin_width=1.0, intensity_scaling=0.5, additional_metadata=[] - ) - spec_tensors, meta_tensors = tensorize_spectra([spectrum, spectrum], settings) - - assert meta_tensors.shape == Size([2, 0]) - assert spec_tensors.shape == Size([2, 990]) - assert spec_tensors[0, 0] == spec_tensors[0, 490] == 0.5**0.5 - assert spec_tensors[0, -1] == 1 - - -@pytest.fixture() -def dummy_data_generator(): - spectrums = create_test_spectra(4, 3) - selected_pairs = SpectrumPairGenerator( - [ - ("CCCCCCCCCCCCCC", "DDDDDDDDDDDDDD", 0.25), - ("BBBBBBBBBBBBBB", "DDDDDDDDDDDDDD", 0.6666667), - ("AAAAAAAAAAAAAA", "CCCCCCCCCCCCCC", 1.0), - ("AAAAAAAAAAAAAA", "BBBBBBBBBBBBBB", 0.33333334), - ], - spectrums, - True, - 0, - ) - batch_size = 2 - settings = SettingsMS2Deepscore( +def _training_settings(**overrides): + settings = dict( min_mz=10, max_mz=1000, mz_bin_width=0.1, intensity_scaling=0.5, additional_metadata=[], - same_prob_bins=np.array([(-0.01, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)]), - batch_size=batch_size, + same_prob_bins=np.array( + [(-0.01, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)] + ), + batch_size=2, num_turns=4, augment_removal_max=0.0, augment_removal_intensity=0.0, augment_intensity=0.0, augment_noise_max=0, average_inchikey_sampling_count=2, + random_seed=0, ) - return TrainingBatchGenerator(selected_pairs, settings) + settings.update(overrides) + return SettingsMS2Deepscore(**settings) -def test_correct_batch_format_data_generator(dummy_data_generator): - def check_correct_batch_format(batch, batch_size=2): - """Checks that each output has the shape of the batch and the expected tensor shapes""" - assert len(batch) == 5, "expected 5 tensors as output" - for i in range(5): - assert batch[i].shape[0] == batch_size - spec1, spec2, meta1, meta2, targets = batch - assert meta1.shape[1] == meta2.shape[1] == 0 - assert spec1.shape[1] == spec2.shape[1] == 9900 - assert targets.shape[0] == batch_size +def _make_training_batch_generator(): + spectra = create_test_spectra(4, 3) + pair_generator = SpectrumPairGenerator( + list(SELECTED_PAIRS), + spectra, + shuffle=True, + random_seed=0, + ) + return TrainingBatchGenerator(pair_generator, _training_settings()) - batch = dummy_data_generator.__getitem__(0) - check_correct_batch_format(batch) - assert len(dummy_data_generator) == 8 +@pytest.fixture +def dummy_data_generator(): + return _make_training_batch_generator() - for batch in dummy_data_generator: - check_correct_batch_format(batch) - - -def test_equal_sampling_of_spectra(dummy_data_generator): - """Tests that all unique spectra are at least sampled once. - The sampling is random, but for enough repetitions very likely to always happen. - This test is mostly to make sure we don't accidentally implement something - where we just resample the same spectrum every time for one inchikey""" - spectrums = create_test_spectra(4, 3) # the same spectra used for the dummy_data_generator - - tensorized_spectra = [] - epochs = 20 - for _ in range(epochs): - for batch in dummy_data_generator: - for i in range(batch[0].shape[0]): - tensorized_spectra.append(tuple(batch[0][i].tolist())) - tensorized_spectra.append(tuple(batch[1][i].tolist())) - # Count occurrences of each unique tensor, the dummy spectra are generated, so they all result in unique tensors. - tensor_counts = {} - for spectrum_tensor in tensorized_spectra: - if spectrum_tensor in tensor_counts: - tensor_counts[spectrum_tensor] += 1 - else: - tensor_counts[spectrum_tensor] = 1 - # test if all spectra are sampled (at least once) - unique_tensors = tensor_counts.keys() - # Test that each spectrum is sampled. This is not really always true, since we randomly sample spectra, - # but since we sample 640 spectra from 24 options, it is very unlikely (1 in 28 billion) - # that this will result in not sampling all at least once. - # Because we have a fixed seed, this should not result in random failing tests. - assert len(unique_tensors) == 12, "Not all spectra are selected at least once" - - def reverse_tensorize(tensor, list_of_spectra, settings): - """Finds the spectrum in a list of spectra based on the tensorized vesion""" - # Create tensors of the available spectra, to later make it possible to link spectra back to inchikeys again. - tensorized_spectra, _ = tensorize_spectra(list_of_spectra, settings) - list_of_spectrum_tensors = [tuple(tensor.tolist()) for tensor in tensorized_spectra] - assert len(set(list_of_spectrum_tensors)) == len(list_of_spectrum_tensors), ( - "There are repeating tensors, meaning that there are spectra that result in exactly the same tensor. " - "Change the dummy spectra to have unique spectra." - ) - for i, tensorized_spectrum in enumerate(list_of_spectrum_tensors): - if tensorized_spectrum == tensor: - return list_of_spectra[i] - - # get spectrum counts per inchikey (by reverse engineering which tensors belong to which spectrum) - inchikey_counts = Counter() - for unique_tensor, count in tensor_counts.items(): - spectrum = reverse_tensorize(unique_tensor, spectrums, dummy_data_generator.model_settings) - - inchikey = spectrum.get("inchikey")[:14] # pyright: ignore[reportOptionalMemberAccess] - inchikey_counts[inchikey] += count - # Test that the inchikeys are sampled equally - assert max(inchikey_counts.values()) - min(inchikey_counts.values()) < 2 - - -def test_create_data_generator(): - """tests if a the function create_data_generator creates a datagenerator that samples all input spectra - correct distributions of inchikeys and scores are tested in other tests""" - test_spectra = create_test_spectra(8, 3) - settings = SettingsMS2Deepscore( + +def _assert_training_batch(batch, batch_size=2, n_bins=9900): + assert len(batch) == 5 + spec1, spec2, meta1, meta2, targets = batch + + assert spec1.shape == (batch_size, n_bins) + assert spec2.shape == (batch_size, n_bins) + assert meta1.shape == (batch_size, 0) + assert meta2.shape == (batch_size, 0) + assert targets.shape == (batch_size,) + assert targets.dtype == torch.float32 + + +def _make_pos_neg_spectra(): + spectra = create_test_spectra(20, 2) + positive = spectra[:20] + negative = spectra[20:] + for spectrum in positive: + spectrum.set("ionmode", "positive") + for spectrum in negative: + spectrum.set("ionmode", "negative") + return positive, negative + + +def _cross_ionmode_settings(): + return SettingsMS2Deepscore( min_mz=10, max_mz=1000, mz_bin_width=0.1, intensity_scaling=0.5, additional_metadata=[], - same_prob_bins=np.array([(-0.01, 0.75), (0.75, 1)]), + same_prob_bins=np.array([(-0.01, 0.6), (0.6, 1.0)]), + max_inchikey_sampling=300, batch_size=2, num_turns=4, - augment_removal_max=0.0, - augment_removal_intensity=0.0, - augment_intensity=0.0, - augment_noise_max=0, - ) - spectrum_pair_generator = create_spectrum_pair_generator(test_spectra, settings=settings) - data_generator = TrainingBatchGenerator(spectrum_pair_generator=spectrum_pair_generator, settings=settings) - tensorized_spectra = [] - epochs = 20 - for _ in range(epochs): - for batch in data_generator: - for i in range(batch[0].shape[0]): - tensorized_spectra.append(tuple(batch[0][i].tolist())) - tensorized_spectra.append(tuple(batch[1][i].tolist())) - # Count occurrences of each unique tensor, the dummy spectra are generated, so they all result in unique tensors. - tensor_counts = {} - for spectrum_tensor in tensorized_spectra: - if spectrum_tensor in tensor_counts: - tensor_counts[spectrum_tensor] += 1 - else: - tensor_counts[spectrum_tensor] = 1 - # test if all spectra are sampled (at least once) - unique_tensors = tensor_counts.keys() - # Test that each spectrum is sampled. This is not really always true, since we randomly sample spectra, - # but since we sample 640 spectra from 24 options, it is very unlikely (1 in 28 billion) - # that this will result in not sampling all at least once. - # Because we have a fixed seed, this should not result in random failing tests. - assert len(unique_tensors) == len(test_spectra), "Not all spectra are selected at least once" - - -### Tests for EmbeddingEvaluator data generator -def test_generator_initialization(data_generator_embedding_evaluation): - """ - Test if the data generator initializes correctly. - """ - assert len(data_generator_embedding_evaluation.spectrums) == 2 * 25, "Incorrect number of spectrums" - assert ( - data_generator_embedding_evaluation.batch_size - == data_generator_embedding_evaluation.settings.evaluator_distribution_size - ), "Incorrect batch size" - - -def test_batch_generation(data_generator_embedding_evaluation): - """ - Test if batches generated are correct in structure and size. - """ - tanimoto_scores, ms2ds_scores, embeddings = next(data_generator_embedding_evaluation) - assert tanimoto_scores.shape == ( - data_generator_embedding_evaluation.batch_size, - data_generator_embedding_evaluation.batch_size, - ), "Incorrect shape for tanimoto_scores" - assert ms2ds_scores.shape == ( - data_generator_embedding_evaluation.batch_size, - data_generator_embedding_evaluation.batch_size, - ), "Incorrect shape for ms2ds_scores" - assert embeddings.shape[0] == data_generator_embedding_evaluation.batch_size, "Incorrect batch size in embeddings" - - -def test_epoch_end_functionality(data_generator_embedding_evaluation): - """ - Test if the generator correctly resets and shuffles after an epoch. - """ - initial_indexes = data_generator_embedding_evaluation.indexes.copy() - counter = 0 - for _ in data_generator_embedding_evaluation: - counter += 1 - assert counter == 5 - # 2nd run - for _ in data_generator_embedding_evaluation: - counter += 1 - assert counter == 10 - assert not np.array_equal(data_generator_embedding_evaluation.indexes, initial_indexes), ( - "Indexes not shuffled after epoch end" + random_seed=11, ) -def test_create_data_generator_across_ionmodes(): - """Just a test that is runs, not a test if it is actually well balanced""" - test_spectra = create_test_spectra(20, 2) - pos_spectra = [] - for spectrum in test_spectra[:20]: - spectrum.set("ionmode", "positive") - pos_spectra.append(spectrum) - neg_spectra = [] - for spectrum in test_spectra[20:]: - spectrum.set("ionmode", "negative") - neg_spectra.append(spectrum) +# --------------------------------------------------------------------------- +# Spectrum tensorization +# --------------------------------------------------------------------------- + +def test_tensorize_spectra_without_metadata(): + spectrum = Spectrum( + mz=np.array([10.0, 500.0, 999.9]), + intensities=np.array([0.5, 0.5, 1.0]), + ) settings = SettingsMS2Deepscore( min_mz=10, max_mz=1000, - mz_bin_width=0.1, + mz_bin_width=1.0, intensity_scaling=0.5, additional_metadata=[], - same_prob_bins=np.array([(-0.01, 0.6), (0.6, 1)]), - max_inchikey_sampling=300, - batch_size=2, - num_turns=4, ) - data_generator = create_data_generator_across_ionmodes(pos_spectra + neg_spectra, settings) - for _ in range(len(data_generator)): - spectra_1, spectra_2, meta_1, meta_2, targets = data_generator.__next__() + spec_tensors, meta_tensors = tensorize_spectra([spectrum, spectrum], settings) -def test_select_compound_pairs_wrapper_across_ionmode(): - test_spectra = create_test_spectra(20, 2) - pos_spectra = [] - for spectrum in test_spectra[:20]: - spectrum.set("ionmode", "positive") - pos_spectra.append(spectrum) - neg_spectra = [] - for spectrum in test_spectra[20:]: - spectrum.set("ionmode", "negative") - neg_spectra.append(spectrum) + assert spec_tensors.shape == (2, 990) + assert meta_tensors.shape == (2, 0) + torch.testing.assert_close(spec_tensors[0, 0], torch.tensor(0.5**0.5)) + torch.testing.assert_close(spec_tensors[0, 490], torch.tensor(0.5**0.5)) + torch.testing.assert_close(spec_tensors[0, -1], torch.tensor(1.0)) + + +# --------------------------------------------------------------------------- +# TrainingBatchGenerator / SpectrumPairGenerator +# --------------------------------------------------------------------------- + + +def test_training_batch_generator_batch_contract(dummy_data_generator): + _assert_training_batch(dummy_data_generator[0]) + assert len(dummy_data_generator) == 8 + + batches = list(dummy_data_generator) + assert len(batches) == 8 + for batch in batches: + _assert_training_batch(batch) + + +def test_training_batch_generator_is_reproducible_with_seed(): + generator_1 = _make_training_batch_generator() + generator_2 = _make_training_batch_generator() + + # Augmentation is disabled in this fixture, so the same pair/spectrum RNG seed + # should produce identical batches. + for batch_1, batch_2 in zip(generator_1, generator_2): + for tensor_1, tensor_2 in zip(batch_1, batch_2): + torch.testing.assert_close(tensor_1, tensor_2) + + +def test_spectrum_pair_generator_samples_valid_spectra_and_repeats_indefinitely(): + spectra = create_test_spectra(4, 3) + generator = SpectrumPairGenerator( + list(SELECTED_PAIRS), + spectra, + shuffle=True, + random_seed=0, + ) + + valid_inchikeys = {s.get("inchikey")[:14] for s in spectra} + seen_spectrum_ids = set() + + # Iterate well beyond one pass through selected_inchikey_pairs. This checks both + # cycling and random selection among multiple spectra belonging to one compound. + for _ in range(200): + spectrum_1, spectrum_2, score = next(generator) + assert spectrum_1.get("inchikey")[:14] in valid_inchikeys + assert spectrum_2.get("inchikey")[:14] in valid_inchikeys + assert 0.0 <= float(score) <= 1.0 + seen_spectrum_ids.add(id(spectrum_1)) + seen_spectrum_ids.add(id(spectrum_2)) + + # All four compounds occur in SELECTED_PAIRS, each with three spectra. With the + # fixed RNG seed this is deterministic, while avoiding expensive tensor roundtrips. + assert len(seen_spectrum_ids) == len(spectra) + + +def test_spectrum_pair_generator_has_balanced_compound_frequency_for_fixture(): + generator = SpectrumPairGenerator( + list(SELECTED_PAIRS), + create_test_spectra(4, 3), + shuffle=False, + random_seed=0, + ) + + counts = generator.get_inchikey_counts() + + assert counts == Counter( + { + "AAAAAAAAAAAAAA": 2, + "BBBBBBBBBBBBBB": 2, + "CCCCCCCCCCCCCC": 2, + "DDDDDDDDDDDDDD": 2, + } + ) + + +def test_create_spectrum_pair_generator_returns_pairs_from_input_compounds(): + spectra = create_test_spectra(8, 3) settings = SettingsMS2Deepscore( min_mz=10, max_mz=1000, mz_bin_width=0.1, intensity_scaling=0.5, additional_metadata=[], - same_prob_bins=np.array([(-0.01, 0.6), (0.6, 1)]), - max_inchikey_sampling=300, + same_prob_bins=np.array([(-0.01, 0.75), (0.75, 1.0)]), batch_size=2, num_turns=4, + augment_removal_max=0.0, + augment_removal_intensity=0.0, + augment_intensity=0.0, + augment_noise_max=0, + random_seed=7, ) - spectrum_pair_generator = select_compound_pairs_wrapper_across_ionmode(pos_spectra, neg_spectra, settings) - for _ in range(len(spectrum_pair_generator)): - spectrum_1, spectrum_2, score = spectrum_pair_generator.__next__() + pair_generator = create_spectrum_pair_generator(spectra, settings=settings) + + assert len(pair_generator) > 0 + valid_inchikeys = {s.get("inchikey")[:14] for s in spectra} + for inchikey_1, inchikey_2, score in pair_generator.selected_inchikey_pairs: + assert inchikey_1 in valid_inchikeys + assert inchikey_2 in valid_inchikeys + assert 0.0 <= float(score) <= 1.0 + + # Also ensure the result can actually feed TrainingBatchGenerator. + batch_generator = TrainingBatchGenerator(pair_generator, settings) + _assert_training_batch(next(batch_generator)) + + +# --------------------------------------------------------------------------- +# Embedding-evaluator data generator +# --------------------------------------------------------------------------- + + +def test_embedding_generator_initialization( + data_generator_embedding_evaluation, + mock_ms2ds_model, +): + generator = data_generator_embedding_evaluation + + assert len(generator.spectrums) == 50 + assert generator.batch_size == generator.settings.evaluator_distribution_size == 10 + assert len(generator) == 5 + + # The MS2DeepScore model is used only for inference by this generator. + assert not mock_ms2ds_model.training + assert not mock_ms2ds_model.encoder.training + + +def test_embedding_generator_batch_contract_and_inference_mode( + data_generator_embedding_evaluation, + mock_ms2ds_model, +): + tanimoto_scores, ms2ds_scores, embeddings = next(data_generator_embedding_evaluation) + batch_size = data_generator_embedding_evaluation.batch_size + + assert tanimoto_scores.shape == (batch_size, batch_size) + assert ms2ds_scores.shape == (batch_size, batch_size) + assert embeddings.shape == (batch_size, 128) + assert not embeddings.requires_grad + + # Similarity matrices for a set against itself must be symmetric with unit diagonal. + np.testing.assert_allclose(tanimoto_scores, tanimoto_scores.T, atol=1e-7) + np.testing.assert_allclose(ms2ds_scores, ms2ds_scores.T, atol=1e-7) + np.testing.assert_allclose(np.diag(tanimoto_scores), 1.0, atol=1e-7) + np.testing.assert_allclose(np.diag(ms2ds_scores), 1.0, atol=1e-6) + + # These assertions explicitly protect the eval()/no_grad() inference contract. + assert mock_ms2ds_model.encoder.last_training is False + assert mock_ms2ds_model.encoder.last_grad_enabled is False + + +def test_embedding_generator_resets_after_each_epoch(data_generator_embedding_evaluation): + generator = data_generator_embedding_evaluation + initial_indexes = generator.indexes.copy() + + assert len(list(generator)) == len(generator) == 5 + assert generator.current_index == 0 + indexes_after_epoch_1 = generator.indexes.copy() + assert not np.array_equal(indexes_after_epoch_1, initial_indexes) + + assert len(list(generator)) == 5 + assert generator.current_index == 0 + indexes_after_epoch_2 = generator.indexes.copy() + assert not np.array_equal(indexes_after_epoch_2, indexes_after_epoch_1) + + +def test_embedding_generator_shuffle_is_reproducible_for_same_seed(mock_ms2ds_model): + settings_1 = SettingsEmbeddingEvaluator(evaluator_distribution_size=10, random_seed=77) + settings_2 = SettingsEmbeddingEvaluator(evaluator_distribution_size=10, random_seed=77) + spectra = create_test_spectra(25, 2) + + # Use separate model instances because generator initialization switches them to eval mode. + model_type = type(mock_ms2ds_model) + generator_1 = DataGeneratorEmbeddingEvaluation(spectra, model_type(), settings_1, device="cpu") + generator_2 = DataGeneratorEmbeddingEvaluation(spectra, model_type(), settings_2, device="cpu") + + np.testing.assert_array_equal(generator_1.indexes, generator_2.indexes) + generator_1.on_epoch_end() + generator_2.on_epoch_end() + np.testing.assert_array_equal(generator_1.indexes, generator_2.indexes) + + +# --------------------------------------------------------------------------- +# Cross-ionmode generators +# --------------------------------------------------------------------------- + + +def test_create_data_generator_across_ionmodes_returns_valid_batches(): + positive, negative = _make_pos_neg_spectra() + generator = create_data_generator_across_ionmodes( + positive + negative, + _cross_ionmode_settings(), + ) + + assert len(generator) > 0 + # CombinedSpectrumGenerator alternates same-mode and cross-mode pair sources. + for _ in range(min(len(generator), 6)): + batch = next(generator) + _assert_training_batch(batch) + assert torch.isfinite(batch[-1]).all() + + +def test_cross_ionmode_pair_generator_preserves_ionmode_direction(): + positive, negative = _make_pos_neg_spectra() + pair_generator = select_compound_pairs_wrapper_across_ionmode( + positive, + negative, + _cross_ionmode_settings(), + ) + + assert len(pair_generator) > 0 + for _ in range(len(pair_generator)): + spectrum_1, spectrum_2, score = next(pair_generator) assert spectrum_1.get("ionmode") == "positive" assert spectrum_2.get("ionmode") == "negative" - # it should be an infinite generator, so it should continue after a loop - spectrum_pair_generator.__next__() + assert 0.0 <= float(score) <= 1.0 + + # It is intentionally cyclic/infinite rather than exhausted after one schedule. + spectrum_1, spectrum_2, _ = next(pair_generator) + assert spectrum_1.get("ionmode") == "positive" + assert spectrum_2.get("ionmode") == "negative" diff --git a/tests/test_embedding_evaluator.py b/tests/test_embedding_evaluator.py index 0cced41d..9233ba7a 100644 --- a/tests/test_embedding_evaluator.py +++ b/tests/test_embedding_evaluator.py @@ -1,133 +1,259 @@ -import os -import pytest +import re + import numpy as np +import pytest +import torch from sklearn.datasets import make_regression -from torch import randn -from ms2deepscore.models import EmbeddingEvaluationModel, LinearModel -from ms2deepscore.models import load_linear_model, load_embedding_evaluator -from tests.test_data_generators import data_generator_embedding_evaluation, MockMS2DSModel -from tests.create_test_spectra import create_test_spectra -from ms2deepscore.SettingsMS2Deepscore import SettingsEmbeddingEvaluator - -# This is just to make the ruff linter happy -fixtures = [data_generator_embedding_evaluation] +from ms2deepscore.SettingsMS2Deepscore import SettingsEmbeddingEvaluator +from ms2deepscore.models import ( + EmbeddingEvaluationModel, + LinearModel, + load_embedding_evaluator, + load_linear_model, +) +from tests.create_test_spectra import create_test_spectra @pytest.fixture -def mock_settings(): +def evaluator_settings(): + # Keep the test model intentionally small: architecture semantics are the same, + # while forward/training tests remain fast. return SettingsEmbeddingEvaluator( - evaluator_num_filters=32, - evaluator_depth=6, - evaluator_kernel_size=40, - mini_batch_size=10, - batches_per_iteration=5, + evaluator_num_filters=8, + evaluator_depth=3, + evaluator_kernel_size=15, + mini_batch_size=5, + batches_per_iteration=1, learning_rate=0.001, num_epochs=1, evaluator_distribution_size=10, + random_seed=13, ) @pytest.fixture -def embedding_model(mock_settings): - return EmbeddingEvaluationModel(settings=mock_settings) - - -def test_model_initialization(embedding_model): - """ - Test if the model initializes with the correct number of filters, depth, and kernel size. - """ - assert embedding_model.settings.evaluator_num_filters == 32, "Incorrect number of filters" - assert embedding_model.settings.evaluator_depth == 6, "Incorrect depth" - assert embedding_model.settings.evaluator_kernel_size == 40, "Incorrect kernel size" - - -def test_forward_pass(embedding_model): - """ - Test the forward pass of the model with a mock input. - """ - mock_input = randn(1, 1, 500) - output = embedding_model(mock_input) - assert output.shape == (1, 1), "Output shape is incorrect" - - -def test_model_with_different_input_sizes(embedding_model): - """ - Test the model with different input sizes to ensure it can handle variable sequence lengths. - """ - sizes = [100, 250, 500, 750] - for size in sizes: - mock_input = randn(1, 1, size) +def embedding_model(evaluator_settings): + return EmbeddingEvaluationModel(settings=evaluator_settings) + + +def test_model_initialization_matches_settings(embedding_model, evaluator_settings): + assert embedding_model.settings.get_dict() == evaluator_settings.get_dict() + assert len(embedding_model.inception_block.inception_modules) == evaluator_settings.evaluator_depth + assert embedding_model.fc.in_features == evaluator_settings.evaluator_num_filters * 4 + assert embedding_model.fc.out_features == 1 + + +@pytest.mark.parametrize( + ("batch_size", "embedding_size"), + [(1, 100), (2, 250), (10, 500), (3, 750)], +) +def test_forward_pass_supports_batch_and_embedding_sizes( + embedding_model, + batch_size, + embedding_size, +): + embedding_model.eval() + mock_input = torch.randn(batch_size, 1, embedding_size) + + with torch.no_grad(): output = embedding_model(mock_input) - assert output.shape == (1, 1), f"Output shape is incorrect for input size {size}" + assert output.shape == (batch_size, 1) + assert torch.isfinite(output).all() -def test_model_with_batch_sizes(embedding_model): - """ - Test the model with different input sizes to ensure it can handle variable sequence lengths. - """ - batch_sizes = [1, 2, 10] - for size in batch_sizes: - mock_input = randn(size, 1, 100) - output = embedding_model(mock_input) - assert output.shape == (size, 1), f"Output shape is incorrect for batch size {size}" +def test_compute_embedding_evaluations_returns_flat_numpy_array(embedding_model): + embedding_model.eval() + embeddings = np.random.default_rng(4).normal(size=(7, 128)).astype(np.float32) -def test_model_save_load(tmp_path, embedding_model): - # Save the model - filepath = tmp_path / "embedding_model.pth" - embedding_model.save(filepath) + result = embedding_model.compute_embedding_evaluations(embeddings, device="cpu") + + assert isinstance(result, np.ndarray) + assert result.shape == (7,) + assert np.isfinite(result).all() + + +def test_model_save_load_roundtrip_preserves_settings_weights_and_predictions( + tmp_path, + embedding_model, +): + filepath = tmp_path / "embedding_model.pt" + embedding_model.eval() + test_input = torch.randn(4, 1, 128) + + with torch.no_grad(): + expected_predictions = embedding_model(test_input).clone() - # Load the model + embedding_model.save(filepath) loaded_model = load_embedding_evaluator(filepath) - # Verify if the saved settings and state dict match the original model - assert loaded_model.settings.evaluator_num_filters == embedding_model.settings.evaluator_num_filters + assert loaded_model.settings.get_dict() == embedding_model.settings.get_dict() assert loaded_model.state_dict().keys() == embedding_model.state_dict().keys() + for name, expected in embedding_model.state_dict().items(): + torch.testing.assert_close(loaded_model.state_dict()[name], expected) + + with torch.no_grad(): + actual_predictions = loaded_model(test_input) + torch.testing.assert_close(actual_predictions, expected_predictions) + + +def test_train_embedding_evaluator_updates_model_parameters( + monkeypatch, + embedding_model, + mock_ms2ds_model, +): + class FakeDataGenerator: + def __init__(self, spectrums, ms2ds_model, settings, device="cpu"): + self.batch_size = settings.evaluator_distribution_size + self.embedding_dim = ms2ds_model.model_settings.embedding_dim + + def __iter__(self): + # Zero target MSE, deterministic non-trivial embeddings. + tanimoto_scores = torch.zeros((self.batch_size, self.batch_size)) + ms2ds_scores = torch.zeros_like(tanimoto_scores) + embeddings = torch.linspace( + 0.0, + 1.0, + self.batch_size * self.embedding_dim, + ).reshape(self.batch_size, self.embedding_dim) + yield tanimoto_scores, ms2ds_scores, embeddings + + monkeypatch.setattr( + "ms2deepscore.models.EmbeddingEvaluatorModel.DataGeneratorEmbeddingEvaluation", + FakeDataGenerator, + ) + monkeypatch.setattr( + "ms2deepscore.models.EmbeddingEvaluatorModel.initialize_device", + lambda: torch.device("cpu"), + ) + before = { + name: parameter.detach().clone() + for name, parameter in embedding_model.named_parameters() + } -def test_train_embedding_evaluator(embedding_model, data_generator_embedding_evaluation): - embedding_model.train_evaluator(create_test_spectra(25), MockMS2DSModel()) - embedding = data_generator_embedding_evaluation.__next__()[2] - result = embedding_model.compute_embedding_evaluations(embedding) - assert result.shape == (10,) - + embedding_model.train_evaluator( + training_spectra=create_test_spectra(2), + ms2ds_model=mock_ms2ds_model, + ) -def test_linear_model_fit_predict(): - # Generate a simple regression problem - X, y = make_regression(n_samples=100, n_features=2, noise=0.1) + changed = [ + not torch.equal(before[name], parameter.detach()) + for name, parameter in embedding_model.named_parameters() + ] + assert any(changed), "Training completed without updating any model parameter." + + +def test_train_embedding_evaluator_reports_actual_validation_loss( + monkeypatch, + capsys, + embedding_model, + mock_ms2ds_model, +): + """Regression test for accidentally reporting the last training loss as val loss.""" + + class FakeDataGenerator: + instance_count = 0 + + def __init__(self, spectrums, ms2ds_model, settings, device="cpu"): + self.batch_size = settings.evaluator_distribution_size + self.embedding_dim = ms2ds_model.model_settings.embedding_dim + self.is_validation = FakeDataGenerator.instance_count == 1 + FakeDataGenerator.instance_count += 1 + + def __iter__(self): + embeddings = torch.linspace( + 0.0, + 1.0, + self.batch_size * self.embedding_dim, + ).reshape(self.batch_size, self.embedding_dim) + ms2ds_scores = torch.zeros((self.batch_size, self.batch_size)) + if self.is_validation: + # Deliberately make the validation target very different from training. + tanimoto_scores = torch.full_like(ms2ds_scores, 10.0) + else: + tanimoto_scores = torch.zeros_like(ms2ds_scores) + yield tanimoto_scores, ms2ds_scores, embeddings + + class RecordingMSELoss(torch.nn.Module): + def __init__(self): + super().__init__() + self.values = [] + + def forward(self, outputs, targets): + loss = ((outputs - targets) ** 2).mean() + self.values.append(float(loss.detach())) + return loss + + criterion = RecordingMSELoss() + monkeypatch.setattr( + "ms2deepscore.models.EmbeddingEvaluatorModel.DataGeneratorEmbeddingEvaluation", + FakeDataGenerator, + ) + monkeypatch.setattr( + "ms2deepscore.models.EmbeddingEvaluatorModel.initialize_device", + lambda: torch.device("cpu"), + ) + monkeypatch.setattr( + "ms2deepscore.models.EmbeddingEvaluatorModel.nn.MSELoss", + lambda: criterion, + ) - # Initialize and fit the model - model = LinearModel(degree=2) - model.fit(X, y) + training_spectra = create_test_spectra(2) + validation_spectra = create_test_spectra(2) + embedding_model.train_evaluator( + training_spectra=training_spectra, + validation_spectra=validation_spectra, + ms2ds_model=mock_ms2ds_model, + ) - # Make predictions - predictions = model.predict(X) + stdout = capsys.readouterr().out + match = re.search(r"Val_loss:\s*([0-9.eE+-]+)", stdout) + assert match is not None, "Expected validation loss to be reported." + reported_validation_loss = float(match.group(1)) - # Check predictions shape - assert predictions.shape == y.shape, "Prediction shape mismatch." + # The final criterion call is the validation batch. This catches code such as + # `val_losses.append(loss_value)` where loss_value still refers to training. + expected_validation_loss = criterion.values[-1] + assert reported_validation_loss == pytest.approx(expected_validation_loss, rel=1e-5, abs=1e-6) + # train_evaluator should restore training mode after temporary validation eval mode. + assert embedding_model.training -def test_linear_model_save_load(tmp_path): - temp_filepath = os.path.join(tmp_path, "temp_model.json") - # Generate a simple regression problem - X, y = make_regression(n_samples=100, n_features=2, noise=0.1) +def test_linear_model_fit_predict(): + X, y = make_regression( + n_samples=100, + n_features=2, + noise=0.1, + random_state=7, + ) + model = LinearModel(degree=2) - # Initialize and fit the model - model = LinearModel(degree=3) model.fit(X, y) + predictions = model.predict(X) - # Save the model - model.save(temp_filepath) + assert predictions.shape == y.shape + assert np.isfinite(predictions).all() - # Ensure the file was created - assert os.path.exists(temp_filepath), "Model file was not created." - # Load the model - loaded_model = load_linear_model(temp_filepath) +def test_linear_model_save_load_roundtrip_preserves_predictions(tmp_path): + X, y = make_regression( + n_samples=100, + n_features=2, + noise=0.1, + random_state=11, + ) + model = LinearModel(degree=3) + model.fit(X, y) + expected = model.predict(X) + + filepath = tmp_path / "linear_model.json" + model.save(filepath) + loaded_model = load_linear_model(filepath) + actual = loaded_model.predict(X) - # Verify the loaded model's parameters match the original model's parameters - 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." + assert filepath.is_file() + assert loaded_model.degree == model.degree == 3 + np.testing.assert_allclose(actual, expected, rtol=1e-12, atol=1e-12) diff --git a/tests/test_ms2deepscore_evaluated.py b/tests/test_ms2deepscore_evaluated.py index c2d32cae..850a3965 100644 --- a/tests/test_ms2deepscore_evaluated.py +++ b/tests/test_ms2deepscore_evaluated.py @@ -50,12 +50,15 @@ def test_MS2DeepScore_score_matrix(): spectrums, similarity_measure = get_test_ms2deepscore_evaluated_instance() scores = similarity_measure.matrix(spectrums[:3], spectrums[:4]) if MATCHMS_V1_API: + assert scores.to_array("predicted_absolute_error").shape == (3, 4) scores = scores.to_array("score") + else: + assert scores["predicted_absolute_error"].shape == (3, 4) + scores = scores["score"] expected_scores = np.array([ [1. , 0.9903664 , 0.9908498 , 0.98811793], [0.9903664 , 1. , 0.99399304, 0.9643621 ], [0.9908498 , 0.99399304, 1. , 0.97351074] ]) - assert np.allclose(expected_scores, scores["score"], atol=1e-6), "Expected different scores." - assert scores["predicted_absolute_error"].shape == (3, 4) + assert np.allclose(expected_scores, scores, atol=1e-6), "Expected different scores."