Files
SciMesh/tests/test_similarity_search.py
T

40 lines
1.4 KiB
Python

from __future__ import annotations
from pathlib import Path
from rdkit import Chem, DataStructs
from scimesh.chemistry.dataset import DatasetStats, find_molecule_by_id, iter_valid_molecules
from scimesh.chemistry.fingerprints import fingerprint
from scimesh.workloads.similarity_search import SimilarityMatch, search_similar
def test_search_matches_full_sorting_and_skips_query_and_invalid(
small_dataset: Path,
) -> None:
query = find_molecule_by_id(small_dataset, "QUERY")
result = search_similar(small_dataset, query, top_k=2)
query_smiles = Chem.MolToSmiles(query.molecule, canonical=True)
expected = []
for record in iter_valid_molecules(small_dataset, DatasetStats()):
if record.molecule_id == query.molecule_id:
continue
if Chem.MolToSmiles(record.molecule, canonical=True) == query_smiles:
continue
expected.append(
SimilarityMatch(
DataStructs.TanimotoSimilarity(
fingerprint(query.molecule), fingerprint(record.molecule)
),
record.molecule_id,
record.smiles,
)
)
assert result.matches == sorted(expected, key=SimilarityMatch.sort_key)[:2]
assert "QUERY" not in {match.molecule_id for match in result.matches}
assert "DUPLICATE" not in {match.molecule_id for match in result.matches}
assert result.stats.invalid == 1
assert result.stats.valid == 5