Refactor into modular molecular workloads
This commit is contained in:
@@ -0,0 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def small_dataset(tmp_path: Path) -> Path:
|
||||
path = tmp_path / "molecules.tsv"
|
||||
path.write_text(
|
||||
"chembl_id\tcanonical_smiles\n"
|
||||
"QUERY\tCCO\n"
|
||||
"ALCOHOL\tCCCO\n"
|
||||
"AMINE\tCCN\n"
|
||||
"BENZENE\tc1ccccc1\n"
|
||||
"BROKEN\tnot-a-smiles\n"
|
||||
"DUPLICATE\tCCO\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
return path
|
||||
@@ -0,0 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from rdkit import DataStructs
|
||||
|
||||
from scimesh.chemistry.dataset import DatasetStats, iter_valid_molecules
|
||||
from scimesh.chemistry.fingerprints import fingerprint
|
||||
from scimesh.workloads.similarity_graph import (
|
||||
SimilarityEdge,
|
||||
build_similarity_graph,
|
||||
write_graph_edges,
|
||||
)
|
||||
|
||||
|
||||
def _brute_force_edges(dataset: Path, threshold: float) -> list[SimilarityEdge]:
|
||||
records = list(iter_valid_molecules(dataset, DatasetStats()))
|
||||
edges = []
|
||||
for left_index, left in enumerate(records):
|
||||
for right in records[left_index + 1 :]:
|
||||
similarity = DataStructs.TanimotoSimilarity(
|
||||
fingerprint(left.molecule), fingerprint(right.molecule)
|
||||
)
|
||||
if similarity >= threshold:
|
||||
edges.append(SimilarityEdge(left.molecule_id, right.molecule_id, similarity))
|
||||
return sorted(edges, key=lambda edge: (edge.source_id, edge.target_id, -edge.similarity))
|
||||
|
||||
|
||||
def test_graph_matches_brute_force_and_has_unique_non_self_edges(
|
||||
small_dataset: Path,
|
||||
) -> None:
|
||||
threshold = 0.15
|
||||
result = build_similarity_graph(small_dataset, threshold, block_size=2)
|
||||
|
||||
assert result.edges == _brute_force_edges(small_dataset, threshold)
|
||||
assert result.checked_pairs == 10
|
||||
edge_pairs = [(edge.source_id, edge.target_id) for edge in result.edges]
|
||||
assert all(source != target for source, target in edge_pairs)
|
||||
assert len(edge_pairs) == len(set(edge_pairs))
|
||||
assert result.stats.invalid == 1
|
||||
|
||||
|
||||
def test_graph_is_block_size_independent_and_deterministic(
|
||||
small_dataset: Path, tmp_path: Path
|
||||
) -> None:
|
||||
first = build_similarity_graph(small_dataset, threshold=0.15, block_size=1)
|
||||
second = build_similarity_graph(small_dataset, threshold=0.15, block_size=3)
|
||||
repeated = build_similarity_graph(small_dataset, threshold=0.15, block_size=3)
|
||||
|
||||
assert first.edges == second.edges == repeated.edges
|
||||
first_path = tmp_path / "first.csv"
|
||||
second_path = tmp_path / "second.csv"
|
||||
write_graph_edges(first_path, first.edges)
|
||||
write_graph_edges(second_path, repeated.edges)
|
||||
assert first_path.read_bytes() == second_path.read_bytes()
|
||||
@@ -0,0 +1,39 @@
|
||||
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
|
||||
Reference in New Issue
Block a user