Refactor into modular molecular workloads
This commit is contained in:
@@ -0,0 +1,13 @@
|
||||
"""Built-in SciMesh workloads."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from scimesh.core.registry import WorkloadRegistry
|
||||
from scimesh.workloads.similarity_graph import SimilarityGraphWorkload
|
||||
from scimesh.workloads.similarity_search import SimilaritySearchWorkload
|
||||
|
||||
|
||||
def register_workloads(registry: WorkloadRegistry) -> None:
|
||||
"""Register built-in workloads in one place, outside the main CLI."""
|
||||
registry.register(SimilaritySearchWorkload())
|
||||
registry.register(SimilarityGraphWorkload())
|
||||
@@ -0,0 +1,185 @@
|
||||
"""Exact sparse molecular similarity graph workload."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from rdkit import DataStructs
|
||||
|
||||
from scimesh.chemistry.dataset import DatasetStats, MoleculeRecord, iter_valid_molecules
|
||||
from scimesh.chemistry.fingerprints import fingerprint
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GraphMolecule:
|
||||
"""A valid record and its fingerprint, built once for graph construction."""
|
||||
|
||||
molecule_id: str
|
||||
fingerprint: Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SimilarityEdge:
|
||||
"""A thresholded, undirected similarity edge represented once as i < j."""
|
||||
|
||||
source_id: str
|
||||
target_id: str
|
||||
similarity: float
|
||||
|
||||
|
||||
@dataclass
|
||||
class GraphResult:
|
||||
"""Edges, dataset statistics, and pair-comparison statistics."""
|
||||
|
||||
edges: list[SimilarityEdge]
|
||||
stats: DatasetStats
|
||||
checked_pairs: int
|
||||
elapsed_seconds: float
|
||||
|
||||
|
||||
def _fingerprinted_molecules(
|
||||
tsv_path: Path, max_rows: int | None
|
||||
) -> tuple[list[GraphMolecule], DatasetStats]:
|
||||
stats = DatasetStats()
|
||||
molecules = [
|
||||
GraphMolecule(record.molecule_id, fingerprint(record.molecule))
|
||||
for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows)
|
||||
]
|
||||
return molecules, stats
|
||||
|
||||
|
||||
def build_similarity_graph(
|
||||
tsv_path: Path,
|
||||
threshold: float,
|
||||
block_size: int,
|
||||
max_rows: int | None = None,
|
||||
progress_every: int = 0,
|
||||
) -> GraphResult:
|
||||
"""Build an exact sparse graph without creating a dense similarity matrix."""
|
||||
if not 0.0 <= threshold <= 1.0:
|
||||
raise ValueError("--threshold must be between 0 and 1")
|
||||
if block_size < 1:
|
||||
raise ValueError("--block-size must be a positive integer")
|
||||
|
||||
molecules, stats = _fingerprinted_molecules(tsv_path, max_rows)
|
||||
edges: list[SimilarityEdge] = []
|
||||
checked_pairs = 0
|
||||
started_at = time.perf_counter()
|
||||
next_report = progress_every
|
||||
|
||||
for left_block_start in range(0, len(molecules), block_size):
|
||||
left_block_end = min(left_block_start + block_size, len(molecules))
|
||||
for right_block_start in range(left_block_start, len(molecules), block_size):
|
||||
right_block_end = min(right_block_start + block_size, len(molecules))
|
||||
same_block = left_block_start == right_block_start
|
||||
for left_index in range(left_block_start, left_block_end):
|
||||
right_start = left_index + 1 if same_block else right_block_start
|
||||
for right_index in range(right_start, right_block_end):
|
||||
checked_pairs += 1
|
||||
similarity = DataStructs.TanimotoSimilarity(
|
||||
molecules[left_index].fingerprint,
|
||||
molecules[right_index].fingerprint,
|
||||
)
|
||||
if similarity >= threshold:
|
||||
edges.append(
|
||||
SimilarityEdge(
|
||||
molecules[left_index].molecule_id,
|
||||
molecules[right_index].molecule_id,
|
||||
similarity,
|
||||
)
|
||||
)
|
||||
if progress_every and checked_pairs >= next_report:
|
||||
elapsed = time.perf_counter() - started_at
|
||||
rate = checked_pairs / elapsed if elapsed else 0.0
|
||||
print(
|
||||
f"Checked {checked_pairs:,} pairs | {len(edges):,} edges | "
|
||||
f"{rate:,.0f} pairs/s | {elapsed:.1f}s elapsed",
|
||||
file=sys.stderr,
|
||||
)
|
||||
next_report += progress_every
|
||||
|
||||
elapsed_seconds = time.perf_counter() - started_at
|
||||
edges.sort(key=lambda edge: (edge.source_id, edge.target_id, -edge.similarity))
|
||||
return GraphResult(edges, stats, checked_pairs, elapsed_seconds)
|
||||
|
||||
|
||||
def write_graph_edges(output_path: Path, edges: list[SimilarityEdge]) -> None:
|
||||
"""Write a deterministic sparse edge list CSV."""
|
||||
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
||||
writer = csv.DictWriter(destination, fieldnames=["source_id", "target_id", "similarity"])
|
||||
writer.writeheader()
|
||||
for edge in edges:
|
||||
writer.writerow(
|
||||
{
|
||||
"source_id": edge.source_id,
|
||||
"target_id": edge.target_id,
|
||||
"similarity": f"{edge.similarity:.6f}",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class SimilarityGraphWorkload:
|
||||
"""CLI adapter for exact block-wise sparse similarity graph construction."""
|
||||
|
||||
name = "similarity-graph"
|
||||
help = "Build an exact sparse graph of thresholded molecular similarities."
|
||||
|
||||
def configure_parser(self, parser: argparse.ArgumentParser) -> None:
|
||||
parser.add_argument("input", type=Path, help="Path to ChEMBL TSV file")
|
||||
parser.add_argument(
|
||||
"--threshold", type=float, required=True,
|
||||
help="Create edges at or above this Tanimoto similarity",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--block-size", type=int, default=1_000,
|
||||
help="Number of molecules per comparison block (default: 1000)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-rows", type=int,
|
||||
help="Read only the first N dataset rows",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--progress-every", type=int, default=100_000,
|
||||
help="Print progress after this many pairs; 0 disables it",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-o", "--output", type=Path, default=Path("similarity_graph.csv"),
|
||||
help="Output edge-list CSV path",
|
||||
)
|
||||
|
||||
def run(self, args: argparse.Namespace) -> int:
|
||||
if args.max_rows is not None and args.max_rows < 1:
|
||||
raise ValueError("--max-rows must be a positive integer")
|
||||
if args.progress_every < 0:
|
||||
raise ValueError("--progress-every cannot be negative")
|
||||
result = build_similarity_graph(
|
||||
args.input,
|
||||
args.threshold,
|
||||
args.block_size,
|
||||
args.max_rows,
|
||||
args.progress_every,
|
||||
)
|
||||
write_graph_edges(args.output, result.edges)
|
||||
rate = (
|
||||
result.checked_pairs / result.elapsed_seconds
|
||||
if result.elapsed_seconds
|
||||
else 0.0
|
||||
)
|
||||
print(
|
||||
f"Valid molecules: {result.stats.valid:,} | invalid SMILES: "
|
||||
f"{result.stats.invalid:,} | scanned rows: {result.stats.scanned:,}"
|
||||
)
|
||||
print(
|
||||
f"Checked pairs: {result.checked_pairs:,} | edges: {len(result.edges):,} | "
|
||||
f"{rate:,.0f} pairs/s | {result.elapsed_seconds:.1f}s elapsed"
|
||||
)
|
||||
if result.stats.stopped_early:
|
||||
print("Stopped early because of --max-rows; graph covers only that subset.")
|
||||
print(f"Saved {len(result.edges)} edges to {args.output}")
|
||||
return 0
|
||||
@@ -0,0 +1,225 @@
|
||||
"""Top-k molecular similarity search workload."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import heapq
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from rdkit import Chem, DataStructs
|
||||
from rdkit.Chem import Draw
|
||||
|
||||
from scimesh.chemistry.dataset import (
|
||||
DatasetStats,
|
||||
MoleculeRecord,
|
||||
find_molecule_by_id,
|
||||
iter_valid_molecules,
|
||||
parse_smiles,
|
||||
)
|
||||
from scimesh.chemistry.fingerprints import fingerprint
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SimilarityMatch:
|
||||
"""A candidate ranked by descending similarity and stable tie-breakers."""
|
||||
|
||||
similarity: float
|
||||
molecule_id: str
|
||||
smiles: str
|
||||
|
||||
def sort_key(self) -> tuple[float, str, str]:
|
||||
return (-self.similarity, self.molecule_id, self.smiles)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _HeapEntry:
|
||||
"""Heap item whose minimum is the worst retained match."""
|
||||
|
||||
match: SimilarityMatch
|
||||
|
||||
def __lt__(self, other: object) -> bool:
|
||||
if not isinstance(other, _HeapEntry):
|
||||
return NotImplemented
|
||||
return self.match.sort_key() > other.match.sort_key()
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResult:
|
||||
"""Results and scan statistics for a similarity search."""
|
||||
|
||||
matches: list[SimilarityMatch]
|
||||
stats: DatasetStats
|
||||
|
||||
|
||||
def search_similar(
|
||||
tsv_path: Path,
|
||||
query: MoleculeRecord,
|
||||
top_k: int,
|
||||
max_rows: int | None = None,
|
||||
progress_every: int = 0,
|
||||
) -> SearchResult:
|
||||
"""Stream top-k matches, retaining only a bounded heap in memory."""
|
||||
if top_k < 1:
|
||||
raise ValueError("--top-k must be a positive integer")
|
||||
query_fingerprint = fingerprint(query.molecule)
|
||||
query_canonical_smiles = Chem.MolToSmiles(query.molecule, canonical=True)
|
||||
stats = DatasetStats()
|
||||
heap: list[_HeapEntry] = []
|
||||
started_at = time.perf_counter()
|
||||
last_report_at = started_at
|
||||
last_report_rows = 0
|
||||
|
||||
for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows):
|
||||
candidate_canonical_smiles = Chem.MolToSmiles(record.molecule, canonical=True)
|
||||
if (
|
||||
record.molecule_id == query.molecule_id
|
||||
or candidate_canonical_smiles == query_canonical_smiles
|
||||
):
|
||||
continue
|
||||
match = SimilarityMatch(
|
||||
DataStructs.TanimotoSimilarity(query_fingerprint, fingerprint(record.molecule)),
|
||||
record.molecule_id,
|
||||
record.smiles,
|
||||
)
|
||||
entry = _HeapEntry(match)
|
||||
if len(heap) < top_k:
|
||||
heapq.heappush(heap, entry)
|
||||
elif match.sort_key() < heap[0].match.sort_key():
|
||||
heapq.heapreplace(heap, entry)
|
||||
|
||||
if progress_every and stats.scanned % progress_every == 0:
|
||||
now = time.perf_counter()
|
||||
interval = now - last_report_at
|
||||
total = now - started_at
|
||||
current_rate = (stats.scanned - last_report_rows) / interval if interval else 0.0
|
||||
average_rate = stats.scanned / total if total else 0.0
|
||||
print(
|
||||
f"Processed {stats.scanned:,} rows | {current_rate:,.0f} rows/s current | "
|
||||
f"{average_rate:,.0f} rows/s average | {total:.1f}s elapsed | "
|
||||
f"{stats.valid:,} valid | {stats.invalid:,} invalid | top {len(heap)}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
last_report_at = now
|
||||
last_report_rows = stats.scanned
|
||||
|
||||
return SearchResult(sorted((entry.match for entry in heap), key=SimilarityMatch.sort_key), stats)
|
||||
|
||||
|
||||
def write_search_results(output_path: Path, matches: list[SimilarityMatch]) -> None:
|
||||
"""Write ranked matches to a deterministic CSV file."""
|
||||
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
||||
writer = csv.DictWriter(
|
||||
destination,
|
||||
fieldnames=["rank", "chembl_id", "canonical_smiles", "similarity"],
|
||||
)
|
||||
writer.writeheader()
|
||||
for rank, match in enumerate(matches, start=1):
|
||||
writer.writerow(
|
||||
{
|
||||
"rank": rank,
|
||||
"chembl_id": match.molecule_id,
|
||||
"canonical_smiles": match.smiles,
|
||||
"similarity": f"{match.similarity:.6f}",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def write_search_images(
|
||||
output_dir: Path, query: MoleculeRecord, matches: list[SimilarityMatch], columns: int
|
||||
) -> tuple[Path, Path]:
|
||||
"""Create PNG depictions for the query molecule and retained matches."""
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
query_path = output_dir / "query.png"
|
||||
Draw.MolToImage(
|
||||
query.molecule, size=(600, 400), legend=f"Query: {query.molecule_id}"
|
||||
).save(query_path)
|
||||
|
||||
candidates_path = output_dir / "top_candidates.png"
|
||||
molecules = [parse_smiles(match.smiles) for match in matches]
|
||||
legends = [
|
||||
f"#{rank} {match.molecule_id}\nTanimoto: {match.similarity:.4f}"
|
||||
for rank, match in enumerate(matches, start=1)
|
||||
]
|
||||
Draw.MolsToGridImage(
|
||||
molecules, molsPerRow=columns, subImgSize=(350, 250), legends=legends
|
||||
).save(candidates_path)
|
||||
return query_path, candidates_path
|
||||
|
||||
|
||||
class SimilaritySearchWorkload:
|
||||
"""CLI adapter for streaming top-k molecular similarity search."""
|
||||
|
||||
name = "similarity-search"
|
||||
help = "Find top-k molecules similar to a query molecule."
|
||||
|
||||
def configure_parser(self, parser: argparse.ArgumentParser) -> None:
|
||||
parser.add_argument("input", type=Path, help="Path to ChEMBL TSV file")
|
||||
query_group = parser.add_mutually_exclusive_group(required=True)
|
||||
query_group.add_argument("--query-id", help="ChEMBL ID of the query molecule")
|
||||
query_group.add_argument("--query-smiles", help="SMILES of the query molecule")
|
||||
parser.add_argument(
|
||||
"--top-k", "--top", dest="top_k", type=int, default=20,
|
||||
help="Number of matches to retain (default: 20)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-o", "--output", type=Path, default=Path("similarity_results.csv"),
|
||||
help="Output CSV path",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--progress-every", type=int, default=100_000,
|
||||
help="Print progress after this many rows; 0 disables it",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-rows", type=int,
|
||||
help="Scan only the first N rows after resolving the query",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--images-dir", type=Path,
|
||||
help="Directory for query and top-candidate PNG images",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image-columns", type=int, default=4,
|
||||
help="Number of molecules per row in the candidate image",
|
||||
)
|
||||
|
||||
def run(self, args: argparse.Namespace) -> int:
|
||||
if args.progress_every < 0:
|
||||
raise ValueError("--progress-every cannot be negative")
|
||||
if args.max_rows is not None and args.max_rows < 1:
|
||||
raise ValueError("--max-rows must be a positive integer")
|
||||
if args.image_columns < 1:
|
||||
raise ValueError("--image-columns must be a positive integer")
|
||||
if args.query_id:
|
||||
query = find_molecule_by_id(args.input, args.query_id)
|
||||
else:
|
||||
molecule = parse_smiles(args.query_smiles)
|
||||
if molecule is None:
|
||||
raise ValueError("--query-smiles is invalid")
|
||||
query = MoleculeRecord("query", args.query_smiles, molecule)
|
||||
|
||||
result = search_similar(
|
||||
args.input, query, args.top_k, args.max_rows, args.progress_every
|
||||
)
|
||||
write_search_results(args.output, result.matches)
|
||||
image_paths: tuple[Path, Path] | None = None
|
||||
if args.images_dir:
|
||||
image_paths = write_search_images(
|
||||
args.images_dir, query, result.matches, args.image_columns
|
||||
)
|
||||
|
||||
print(f"Query {query.molecule_id}: {query.smiles}")
|
||||
print(
|
||||
f"Scanned {result.stats.scanned:,} rows: {result.stats.valid:,} valid, "
|
||||
f"{result.stats.invalid:,} invalid SMILES."
|
||||
)
|
||||
if result.stats.stopped_early:
|
||||
print("Stopped early because of --max-rows; results cover only that subset.")
|
||||
print(f"Saved {len(result.matches)} matches to {args.output}")
|
||||
if image_paths:
|
||||
print(f"Saved query image to {image_paths[0]}")
|
||||
print(f"Saved candidate image to {image_paths[1]}")
|
||||
return 0
|
||||
Reference in New Issue
Block a user