diff --git a/benchmarks/substruct_library_bench.py b/benchmarks/substruct_library_bench.py new file mode 100644 index 00000000..60668e02 --- /dev/null +++ b/benchmarks/substruct_library_bench.py @@ -0,0 +1,403 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Benchmark persistent nvMolKit and RDKit substructure libraries. + +Each backend configuration stages and finalizes one library, timed once, then +times every requested operation against it. nvMolKit results are validated +against RDKit over the queries RDKit completed before its deadline. + +Examples: + python substruct_library_bench.py --smiles molecules.smi --smarts queries.smarts + python substruct_library_bench.py --smiles molecules.smi --smarts queries.smarts \ + --operations has get --algorithms gsi dfs --rdkit_holders mol cached-pattern + python substruct_library_bench.py --pickle targets.pkl --query_smiles queries.smi --num_queries 1000 \ + --query_modes serial concurrent --no-rdkit +""" + +from __future__ import annotations + +import argparse +import random +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Any + +from bench_utils import ( + Deadline, + TimingResult, + add_backend_selection_args, + add_rdkit_max_seconds_arg, + load_pickle, + load_smarts, + load_smiles, + print_csv_rows, + throughput_per_s, + time_it, + write_csv_rows, +) +from rdkit import Chem +from rdkit.Chem import rdSubstructLibrary + +OPERATIONS = ("has", "count", "get") + + +@dataclass(frozen=True) +class Measurement: + """One-time build timings, steady-state query timing, and the last query sweep's results.""" + + staging_ms: float + finalize_ms: float + steady_ms: float + steady_std_ms: float + results: list[Any] + completed_queries: int + + +def run_nvmolkit_queries( + library: Any, queries: Sequence[Any], operation: str, max_results: int, query_mode: str +) -> list[Any]: + """Run every query, waiting on each in turn (serial) or submitting all before waiting (concurrent).""" + if query_mode == "serial": + run = { + "has": library.hasMatchSync, + "count": library.countMatchesSync, + "get": lambda query: library.getMatchesSync(query, max_results), + }[operation] + return [run(query) for query in queries] + submit = { + "has": library.hasMatch, + "count": library.countMatches, + "get": lambda query: library.getMatches(query, max_results), + }[operation] + return [future.result() for future in [submit(query) for query in queries]] + + +def run_rdkit_queries( + library: Any, + queries: Sequence[Any], + operation: str, + max_results: int, + num_threads: int, + deadline: Deadline, +) -> list[Any]: + """Run queries until the deadline expires; results cover a prefix of the query list.""" + options = { + "recursionPossible": True, + "useChirality": False, + "useQueryQueryMatches": False, + "numThreads": num_threads, + } + results = [] + for query in queries: + if deadline.expired(): + break + if operation == "has": + results.append(library.HasMatch(query, **options)) + elif operation == "count": + results.append(library.CountMatches(query, **options)) + else: + results.append(list(library.GetMatches(query, maxResults=max_results, **options))) + return results + + +def time_rdkit_operation( + library: Any, + queries: Sequence[Any], + operation: str, + *, + max_results: int, + num_threads: int, + runs: int, + max_seconds: float, +) -> tuple[TimingResult, list[Any]]: + """Time one operation's query sweep under a shared deadline; return the timing and matching results. + + There is no warmup: time_it runs warmups without the deadline, which would defeat it. The returned results come + from a sweep time_it keeps: a complete one, or the first sweep when even that was cut short. + """ + latest: list[Any] = [] + kept: list[Any] | None = None + + def search(deadline: Deadline) -> None: + nonlocal latest, kept + latest = run_rdkit_queries(library, queries, operation, max_results, num_threads, deadline) + if kept is None or len(latest) == len(queries): + kept = latest + + steady = time_it( + search, + runs=runs, + warmups=0, + max_seconds=max_seconds, + progress_getter=lambda: len(latest), + progress_target=len(queries), + ) + return steady, kept or [] + + +def benchmark_rdkit( + mols: Sequence[Any], + queries: Sequence[Any], + *, + operations: Sequence[str], + holder: str, + num_threads: int, + max_results: int, + runs: int, + max_seconds: float, +) -> dict[str, Measurement]: + """Build one RDKit library and time every operation against it.""" + + def stage() -> None: + for mol in mols: + library.GetMolHolder().AddMol(mol) + if holder == "cached-pattern": + rdSubstructLibrary.AddPatterns(library, numThreads=num_threads) + + mol_holder = rdSubstructLibrary.MolHolder() if holder == "mol" else rdSubstructLibrary.CachedMolHolder() + library = rdSubstructLibrary.SubstructLibrary(mol_holder) + staging = time_it(stage, runs=1, warmups=0) + + measurements = {} + for operation in operations: + steady, results = time_rdkit_operation( + library, + queries, + operation, + max_results=max_results, + num_threads=num_threads, + runs=runs, + max_seconds=max_seconds, + ) + measurements[operation] = Measurement( + staging.mean_ms, 0.0, steady.mean_ms, steady.std_ms, results, steady.progress + ) + return measurements + + +def benchmark_nvmolkit( + mols: Sequence[Any], + queries: Sequence[Any], + *, + operations: Sequence[str], + query_modes: Sequence[str], + config: Any, + max_results: int, + runs: int, + warmups: int, +) -> dict[tuple[str, str], Measurement]: + """Build one nvMolKit library and time every operation and query mode against it.""" + from nvmolkit.substruct_library import SubstructLibrary + + library = SubstructLibrary(config=config) + staging = time_it(lambda: library.addMols(mols), runs=1, warmups=0) + finalization = time_it(library.finalize, runs=1, warmups=0, gpu_sync=True) + + measurements = {} + for operation in operations: + for query_mode in query_modes: + results: list[Any] = [] + + def search(operation: str = operation, query_mode: str = query_mode) -> None: + nonlocal results + results = run_nvmolkit_queries(library, queries, operation, max_results, query_mode) + + steady = time_it(search, runs=runs, warmups=warmups, gpu_sync=True) + measurements[(operation, query_mode)] = Measurement( + staging.mean_ms, finalization.mean_ms, steady.mean_ms, steady.std_ms, results, len(queries) + ) + return measurements + + +def first_mismatch(actual: Sequence[Any], expected: Sequence[Any], label: str) -> str | None: + """Describe the first differing query over the prefix both result lists cover, or return None.""" + for query_index, (got, want) in enumerate(zip(actual, expected)): + if got != want: + return f"{label}: query {query_index} gave {got!r}, RDKit reference {want!r}" + return None + + +def result_row( + measurement: Measurement, *, operation: str, num_mols: int, num_queries: int, **configuration: Any +) -> dict[str, Any]: + """One CSV row: configuration, timings, and steady and amortized throughput. + + Amortized figures charge the library's one-time staging and finalize cost to this row's operation alone. + """ + completed = measurement.completed_queries + pairs = num_mols * completed + amortized_ms = measurement.staging_ms + measurement.finalize_ms + measurement.steady_ms + return { + "operation": operation, + "num_mols": num_mols, + "num_queries": num_queries, + "completed_queries": completed, + "positive_queries": sum(bool(result) for result in measurement.results), + **configuration, + "staging_ms": measurement.staging_ms, + "finalize_ms": measurement.finalize_ms, + "steady_ms": measurement.steady_ms, + "steady_std_ms": measurement.steady_std_ms, + "amortized_ms": amortized_ms, + "steady_queries_per_s": throughput_per_s(completed, measurement.steady_ms), + "steady_pairs_per_s": throughput_per_s(pairs, measurement.steady_ms), + "amortized_queries_per_s": throughput_per_s(completed, amortized_ms), + } + + +def load_queries(args: argparse.Namespace) -> list[Any]: + """Load SMARTS queries, or SMILES queries with stereochemistry removed, sampling num_queries of them.""" + if args.smarts: + queries, _ = load_smarts(args.smarts) + if 0 < args.num_queries < len(queries): + queries = random.Random(args.seed).sample(queries, args.num_queries) + return queries + queries = load_smiles(args.query_smiles, args.num_queries, args.sanitize, seed=args.seed) + for query in queries: + Chem.RemoveStereochemistry(query) + return queries + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description="Persistent SubstructLibrary benchmark: nvMolKit vs RDKit") + targets = parser.add_mutually_exclusive_group(required=True) + targets.add_argument("--smiles", "-s", help="SMILES target file") + targets.add_argument("--pickle", help="Pickled RDKit molecule binaries") + query_inputs = parser.add_mutually_exclusive_group(required=True) + query_inputs.add_argument("--smarts", "-q", help="SMARTS query file") + query_inputs.add_argument("--query_smiles", help="SMILES query file; stereochemistry is ignored") + parser.add_argument("--num_mols", "-n", type=int, default=0, help="Maximum target molecules; 0 means all") + parser.add_argument("--num_queries", type=int, default=0, help="Sample this many queries; 0 means all") + parser.add_argument("--seed", type=int, default=42, help="Sampling seed") + parser.add_argument("--no_sanitize", dest="sanitize", action="store_false") + parser.add_argument("--operations", nargs="+", choices=OPERATIONS, default=["has"]) + parser.add_argument("--max_results", type=int, default=-1, help="Result limit for get; -1 means all") + parser.add_argument("--query_modes", nargs="+", choices=["serial", "concurrent"], default=["serial"]) + parser.add_argument("--algorithms", nargs="+", choices=["gsi", "dfs"], default=["gsi", "dfs"]) + parser.add_argument("--batch_size", type=int, default=1024) + parser.add_argument("--workers", type=int, default=-1) + parser.add_argument("--prep_threads", type=int, default=-1, help="Preprocessing threads (-1 = auto)") + parser.add_argument("--gpu_ids", nargs="+", type=int, default=[0], help="GPUs one library is sharded across") + parser.add_argument("--rdkit_holders", nargs="+", choices=["mol", "cached-pattern"], default=["cached-pattern"]) + parser.add_argument("--rdkit_threads", type=int, default=-1) + add_rdkit_max_seconds_arg(parser, extra_help="The deadline is checked between queries.") + parser.add_argument("--runs", "-r", type=int, default=3) + parser.add_argument( + "--warmups", type=int, default=1, help="nvMolKit warmup sweeps; RDKit runs none so its deadline holds" + ) + parser.add_argument("--no_validate", dest="validate", action="store_false") + parser.add_argument("--output", "-o", help="Optional CSV output path") + add_backend_selection_args(parser) + return parser + + +def main() -> None: + args = build_parser().parse_args() + if args.no_rdkit and args.no_nvmolkit: + raise ValueError("cannot disable both backends") + if args.max_results == 0 or args.max_results < -1: + raise ValueError("max_results must be -1 or positive") + if args.pickle: + mols = load_pickle(args.pickle, args.num_mols, seed=args.seed) + else: + mols = load_smiles(args.smiles, args.num_mols, args.sanitize, seed=args.seed) + queries = load_queries(args) + if not mols or not queries: + raise ValueError("no valid target molecules or queries loaded") + + rows = [] + references: dict[str, list[Any]] = {} + mismatches: list[str] = [] + if not args.no_rdkit: + for holder in args.rdkit_holders: + measurements = benchmark_rdkit( + mols, + queries, + operations=args.operations, + holder=holder, + num_threads=args.rdkit_threads, + max_results=args.max_results, + runs=args.runs, + max_seconds=args.rdkit_max_seconds, + ) + for operation, measurement in measurements.items(): + # RDKit holders must agree with each other before they serve as the reference. + mismatch = first_mismatch( + measurement.results, references.get(operation, []), f"rdkit {holder} {operation}" + ) + if mismatch is not None: + mismatches.append(mismatch) + if len(measurement.results) > len(references.get(operation, [])): + references[operation] = measurement.results + rows.append( + result_row( + measurement, + operation=operation, + num_mols=len(mols), + num_queries=len(queries), + backend="rdkit", + holder=holder, + rdkit_threads=args.rdkit_threads, + max_results=args.max_results, + ) + ) + + if not args.no_nvmolkit: + import torch + + from nvmolkit.substructure import SubstructSearchConfig + + torch.cuda.set_device(args.gpu_ids[0]) + for algorithm in args.algorithms: + config = SubstructSearchConfig( + batchSize=args.batch_size, + workerThreads=args.workers, + preprocessingThreads=args.prep_threads, + gpuIds=args.gpu_ids, + algorithm=algorithm, + ) + measurements = benchmark_nvmolkit( + mols, + queries, + operations=args.operations, + query_modes=args.query_modes, + config=config, + max_results=args.max_results, + runs=args.runs, + warmups=args.warmups, + ) + for (operation, query_mode), measurement in measurements.items(): + label = f"nvmolkit {algorithm} {query_mode} {operation}" + if args.validate: + reference = references.get(operation, []) + print(f"VALIDATION {label}: compared {len(reference)}/{len(queries)} queries", flush=True) + mismatch = first_mismatch(measurement.results, reference, label) + if mismatch is not None: + mismatches.append(mismatch) + rows.append( + result_row( + measurement, + operation=operation, + num_mols=len(mols), + num_queries=len(queries), + backend="nvmolkit", + algorithm=algorithm, + query_mode=query_mode, + num_gpus=len(args.gpu_ids), + gpu_ids=" ".join(str(gpu) for gpu in args.gpu_ids), + batch_size=args.batch_size, + workers=args.workers, + prep_threads=args.prep_threads, + max_results=args.max_results, + ) + ) + + print_csv_rows(rows) + write_csv_rows(rows, args.output) + if mismatches: + raise AssertionError("results differ from RDKit:\n" + "\n".join(mismatches)) + + +if __name__ == "__main__": + main() diff --git a/benchmarks/tests/test_substruct_library_bench.py b/benchmarks/tests/test_substruct_library_bench.py new file mode 100644 index 00000000..8647683e --- /dev/null +++ b/benchmarks/tests/test_substruct_library_bench.py @@ -0,0 +1,122 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import pytest +from bench_utils import Deadline +from rdkit import Chem +from rdkit.Chem import rdSubstructLibrary +from substruct_library_bench import ( + Measurement, + benchmark_nvmolkit, + benchmark_rdkit, + first_mismatch, + load_queries, + result_row, + run_rdkit_queries, +) + +TARGETS = ["CCO", "c1ccccc1O", "CC(=O)O", "CCN", "c1ccncc1", "OCCO"] +QUERIES = ["[OX2H]", "c1ccccc1", "[#7]", "C=O", "[Si]"] + + +@pytest.fixture +def mols(): + return [Chem.MolFromSmiles(smiles) for smiles in TARGETS] + + +@pytest.fixture +def queries(): + return [Chem.MolFromSmarts(smarts) for smarts in QUERIES] + + +def expected_matches(mols, queries): + return [[index for index, mol in enumerate(mols) if mol.HasSubstructMatch(query)] for query in queries] + + +@pytest.mark.parametrize("holder", ["mol", "cached-pattern"]) +def test_rdkit_backend_matches_direct_substructure_search(mols, queries, holder): + measurements = benchmark_rdkit( + mols, + queries, + operations=["has", "count", "get"], + holder=holder, + num_threads=1, + max_results=-1, + runs=1, + max_seconds=0, + ) + expected = expected_matches(mols, queries) + assert measurements["get"].results == expected + assert measurements["count"].results == [len(matches) for matches in expected] + assert measurements["has"].results == [bool(matches) for matches in expected] + assert all(measurement.completed_queries == len(queries) for measurement in measurements.values()) + + +def test_rdkit_queries_stop_at_the_deadline(mols, queries): + library = rdSubstructLibrary.SubstructLibrary(rdSubstructLibrary.MolHolder()) + for mol in mols: + library.AddMol(mol) + + class ExpiresAfter(Deadline): + def __init__(self, checks): + super().__init__(0) + self.checks = checks + + def expired(self): + self.checks -= 1 + return self.checks < 0 + + assert run_rdkit_queries(library, queries, "get", -1, 1, ExpiresAfter(2)) == expected_matches(mols, queries)[:2] + + +def test_validation_compares_the_completed_prefix(): + assert first_mismatch([[0, 2], [1], [3]], [[0, 2], [1]], "get") is None + assert "query 1" in first_mismatch([[0, 2], [1]], [[0, 2], [1, 4]], "get") + + +def test_result_row_reports_throughput_over_completed_queries(): + measurement = Measurement( + staging_ms=10, finalize_ms=20, steady_ms=50, steady_std_ms=1, results=[[1], [], [2, 3]], completed_queries=3 + ) + row = result_row(measurement, operation="get", num_mols=100, num_queries=4, backend="nvmolkit") + + assert row["num_queries"] == 4 + assert row["positive_queries"] == 2 + assert row["amortized_ms"] == 80 + assert row["steady_queries_per_s"] == pytest.approx(60) + assert row["steady_pairs_per_s"] == pytest.approx(6000) + assert row["amortized_queries_per_s"] == pytest.approx(37.5) + + +def test_smiles_queries_ignore_stereochemistry(tmp_path): + path = tmp_path / "queries.smi" + path.write_text("C[C@H](O)CC\n") + args = SimpleNamespace(smarts=None, query_smiles=str(path), num_queries=0, sanitize=True, seed=0) + + (query,) = load_queries(args) + assert Chem.MolToSmiles(query) == "CCC(C)O" + + +def test_nvmolkit_backend_matches_rdkit_in_every_query_mode(mols, queries): + torch = pytest.importorskip("torch") + if not torch.cuda.is_available(): + pytest.skip("requires a CUDA device") + from nvmolkit.substructure import SubstructSearchConfig + + measurements = benchmark_nvmolkit( + mols, + queries, + operations=["has", "count", "get"], + query_modes=["serial", "concurrent"], + config=SubstructSearchConfig(), + max_results=-1, + runs=1, + warmups=0, + ) + expected = expected_matches(mols, queries) + for mode in ("serial", "concurrent"): + assert measurements[("get", mode)].results == expected + assert measurements[("count", mode)].results == [len(matches) for matches in expected] + assert measurements[("has", mode)].results == [bool(matches) for matches in expected] diff --git a/docs/api/nvmolkit.rst b/docs/api/nvmolkit.rst index 037ce4a1..107a7c8c 100644 --- a/docs/api/nvmolkit.rst +++ b/docs/api/nvmolkit.rst @@ -146,6 +146,7 @@ Substructure Search substructure.SubstructSearchConfig substructure.SubstructMatchResults + substruct_library.SubstructLibrary Conformer RMSD -------------- diff --git a/docs/conf.py b/docs/conf.py index 1756e1e4..b533714b 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -108,6 +108,7 @@ def __getattr__(self, name): "_Fingerprints", "_mcs", "_mmffOptimization", + "_substructLibrary", "_substructure", "_TFD", "_types", diff --git a/docs/index.rst b/docs/index.rst index 94712a8c..9c5cbafa 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -186,6 +186,7 @@ nvMolKit currently supports the following features: * **Substructure Search**: GPU-accelerated substructure matching against batches of molecules * Supports SMILES and recursive SMARTS-based query molecules via RDKit + * ``SubstructLibrary`` keeps target molecules on one or more GPUs for repeated queries, as a GPU counterpart to RDKit's ``SubstructLibrary`` * Does not yet support chirality-aware matching, enhanced stereochemistry, or other advanced RDKit ``SubstructMatchParameters`` options * **Maximum Common Substructure (MCS)**: GPU-accelerated MCS search across batches of molecule pairs diff --git a/nvmolkit/CMakeLists.txt b/nvmolkit/CMakeLists.txt index eec68417..79c437c6 100644 --- a/nvmolkit/CMakeLists.txt +++ b/nvmolkit/CMakeLists.txt @@ -138,6 +138,14 @@ target_include_directories(_substructure PUBLIC ${Boost_INCLUDE_DIRS} ${Python_INCLUDE_DIRS}) installpythontarget(_substructure ./) +add_library(_substructLibrary MODULE substructLibrary.cpp) +target_link_libraries(_substructLibrary PUBLIC ${Boost_LIBRARIES} + ${PYTHON_LIBRARIES}) +target_link_libraries(_substructLibrary PRIVATE ${RDKit_LIBS} substruct_library) +target_include_directories(_substructLibrary PUBLIC ${Boost_INCLUDE_DIRS} + ${Python_INCLUDE_DIRS}) +installpythontarget(_substructLibrary ./) + add_library(_mcs MODULE mcs.cpp) target_link_libraries(_mcs PUBLIC ${Boost_LIBRARIES} ${PYTHON_LIBRARIES}) target_link_libraries(_mcs PRIVATE ${RDKit_LIBS} mcs_search CUDA::cudart) diff --git a/nvmolkit/substructLibrary.cpp b/nvmolkit/substructLibrary.cpp new file mode 100644 index 00000000..674051e8 --- /dev/null +++ b/nvmolkit/substructLibrary.cpp @@ -0,0 +1,109 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +#include + +#include +#include +#include +#include +#include + +#include "src/substruct/substruct_library.h" + +namespace { + +using namespace boost::python; + +class ScopedGilRelease { + public: + ScopedGilRelease() : state_(PyEval_SaveThread()) {} + ScopedGilRelease(const ScopedGilRelease&) = delete; + ScopedGilRelease& operator=(const ScopedGilRelease&) = delete; + ~ScopedGilRelease() { PyEval_RestoreThread(state_); } + + private: + PyThreadState* state_; +}; + +unsigned int addMolecule(nvMolKit::SubstructLibrary& library, const RDKit::ROMol& molecule) { + const ScopedGilRelease release; + return library.addMol(molecule); +} + +list addMolecules(nvMolKit::SubstructLibrary& library, const object& molecules) { + std::vector owners; + std::vector pointers; + stl_input_iterator iterator(molecules), end; + for (; iterator != end; ++iterator) { + owners.push_back(*iterator); + pointers.push_back(extract(owners.back())); + } + + std::vector ids; + { + const ScopedGilRelease release; + ids = library.addMols(pointers); + } + + list result; + for (const unsigned int id : ids) { + result.append(id); + } + return result; +} + +void finalize(nvMolKit::SubstructLibrary& library) { + const ScopedGilRelease release; + library.finalize(); +} + +list getMatches(nvMolKit::SubstructLibrary& library, const RDKit::ROMol& query, int maxResults) { + std::vector ids; + { + const ScopedGilRelease release; + ids = library.getMatches(query, maxResults); + } + list result; + for (const unsigned int id : ids) { + result.append(id); + } + return result; +} + +std::size_t countMatches(nvMolKit::SubstructLibrary& library, const RDKit::ROMol& query) { + const ScopedGilRelease release; + return library.countMatches(query); +} + +bool hasMatch(nvMolKit::SubstructLibrary& library, const RDKit::ROMol& query) { + const ScopedGilRelease release; + return library.hasMatch(query); +} + +// Accessors wait on the library's lock, so they release the GIL like the operations do. +template +std::size_t withoutGil(const nvMolKit::SubstructLibrary& library) { + const ScopedGilRelease release; + return (library.*Accessor)(); +} + +} // namespace + +BOOST_PYTHON_MODULE(_substructLibrary) { + import("rdkit.Chem.rdchem"); + import("nvmolkit._substructure"); + + class_( + "SubstructLibrary", + init((arg("config") = nvMolKit::SubstructSearchConfig()))) + .def("addMol", &addMolecule) + .def("addMols", &addMolecules) + .def("finalize", &finalize) + .def("getMatches", &getMatches, (arg("query"), arg("maxResults") = -1)) + .def("countMatches", &countMatches) + .def("hasMatch", &hasMatch) + .def("__len__", &withoutGil<&nvMolKit::SubstructLibrary::size>) + .add_property("pendingSize", &withoutGil<&nvMolKit::SubstructLibrary::pendingSize>) + .add_property("maxConcurrentQueries", &withoutGil<&nvMolKit::SubstructLibrary::maxConcurrentQueries>); +} diff --git a/nvmolkit/substruct_library.py b/nvmolkit/substruct_library.py new file mode 100644 index 00000000..513fe086 --- /dev/null +++ b/nvmolkit/substruct_library.py @@ -0,0 +1,179 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Molecules kept on the GPU for repeated substructure queries.""" + +from __future__ import annotations + +import threading +from collections.abc import Callable, Iterable +from concurrent.futures import Future, ThreadPoolExecutor +from typing import Any, TypeVar + +from rdkit.Chem import Mol + +from nvmolkit._substructLibrary import SubstructLibrary as _NativeSubstructLibrary +from nvmolkit.substructure import SubstructSearchConfig + +__all__ = ["SubstructLibrary"] + +_T = TypeVar("_T") + + +class SubstructLibrary: + """A collection of molecules kept on the GPU for repeated substructure searches. + + This is the GPU counterpart of RDKit's ``rdSubstructLibrary.SubstructLibrary``: + target molecules are prepared and uploaded once, then any number of queries + run against them. Use it when the same targets are searched many times; for a + one-off batch of targets and queries, :func:`nvmolkit.substructure.hasSubstructMatch` + avoids the upload step. + + Molecules added with :meth:`addMol` or :meth:`addMols` receive stable indices + in insertion order and become searchable after :meth:`finalize`. Queries keep + seeing the last finalized collection while newly added molecules are pending, + so the library can grow between query rounds. + + Which calls wait: + + * :meth:`addMol`, :meth:`addMols`, and :meth:`finalize` return when their work + is done. They also wait for any queries currently running to finish. + * :meth:`getMatches`, :meth:`countMatches`, and :meth:`hasMatch` return a + :class:`concurrent.futures.Future` right away and search in the background, + so you can keep working or start more queries. Up to + :attr:`maxConcurrentQueries` queries run at once; the rest wait their turn. + Call ``result()`` on the future to get the answer. + * :meth:`getMatchesSync`, :meth:`countMatchesSync`, and :meth:`hasMatchSync` + start a query the same way and return its answer once it is ready. + + For the best throughput, start all of your queries before calling ``result()`` + on the first one, rather than waiting for each answer before starting the next. + + A query's errors, including querying before the first :meth:`finalize`, are + raised by the future's ``result()`` or by the ``*Sync`` call. A query searches + the molecules that were finalized when it begins running, which may include + molecules finalized after it was started. Do not wait for a query's result + inside another future's done-callback; that can hang. + + The library can be used from several threads at once. + + Matching follows RDKit's ``SubstructLibrary`` with ``useChirality=False``: + recursive SMARTS are supported, and stereochemistry in targets and in SMILES + queries is ignored. SMARTS queries that specify atom chirality (``@`` or + ``@@``) are not supported; their results raise ``RuntimeError``. Targets outside the + GPU representation limits (more than 128 atoms, high-degree atoms, isotopes + above 255, or dative and other unsupported bond types) are matched with + RDKit on the CPU, so results cover every added molecule. + + Example:: + + from rdkit import Chem + from nvmolkit.substruct_library import SubstructLibrary + + library = SubstructLibrary() + library.addMols(Chem.MolFromSmiles(smiles) for smiles in ["CCO", "c1ccccc1O", "CCN"]) + library.finalize() + library.getMatchesSync(Chem.MolFromSmarts("[OX2H]")) # [0, 1] + + # Start several queries, then collect their results. + futures = [library.hasMatch(Chem.MolFromSmarts(smarts)) for smarts in ["N", "S"]] + [future.result() for future in futures] # [True, False] + + Args: + config: Search configuration. ``algorithm`` selects the ``"dfs"`` (default) + or ``"gsi"`` backend, and ``gpuIds`` spreads the molecules across + several GPUs. ``batchSize``, ``workerThreads``, and + ``preprocessingThreads`` tune throughput; ``maxMatches`` and + ``uniquify`` do not apply because queries return molecule indices. + """ + + def __init__(self, config: SubstructSearchConfig | None = None) -> None: + """Create an empty library.""" + if config is None: + config = SubstructSearchConfig() + self._native = _NativeSubstructLibrary(config._as_native()) + self._executorLock = threading.Lock() + self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="nvmolkit-substruct") + + def __len__(self) -> int: + """Return the number of searchable molecules.""" + return len(self._native) + + @property + def pendingSize(self) -> int: + """Number of molecules added since the last :meth:`finalize`.""" + return self._native.pendingSize + + @property + def maxConcurrentQueries(self) -> int: + """How many queries can run at once, based on the GPU memory left after :meth:`finalize`; 0 before it.""" + return int(self._native.maxConcurrentQueries) + + def addMol(self, molecule: Mol) -> int: + """Copy ``molecule`` into the library and return its index. + + The molecule becomes searchable after the next :meth:`finalize`. + """ + return self._native.addMol(molecule) + + def addMols(self, molecules: Iterable[Mol]) -> list[int]: + """Copy ``molecules`` into the library and return their indices.""" + return self._native.addMols(molecules) + + def finalize(self) -> None: + """Upload molecules added since the last call and make them searchable. + + Must be called at least once before querying; until then, query results + raise ``RuntimeError``. Queries that have not begun running yet will search + the newly finalized molecules. If it raises, the previously finalized + molecules stay searchable and the new ones stay pending. + """ + try: + self._native.finalize() + finally: + # A failed finalize keeps the previous molecules searchable, so always size the + # executor to what the native library now allows. Queued queries keep running. + with self._executorLock: + previous = self._executor + self._executor = ThreadPoolExecutor( + max_workers=max(1, self.maxConcurrentQueries), + thread_name_prefix="nvmolkit-substruct", + ) + previous.shutdown(wait=False) + + def _submit(self, function: Callable[..., _T], *args: Any) -> Future[_T]: + with self._executorLock: + return self._executor.submit(function, *args) + + def getMatches(self, query: Mol, maxResults: int = -1) -> Future[list[int]]: + """Start a query and return a future for the indices of matching molecules, without waiting. + + Args: + query: Query molecule, typically from ``Chem.MolFromSmarts`` or ``Chem.MolFromSmiles``. + maxResults: Return at most this many indices, keeping the lowest. -1 (default) returns all. + + Returns: + Future resolving to matching indices in ascending order. Its ``result()`` + raises any error from the search. + """ + return self._submit(self._native.getMatches, query, int(maxResults)) + + def countMatches(self, query: Mol) -> Future[int]: + """Start a query and return a future for the number of matching molecules, without waiting.""" + return self._submit(self._native.countMatches, query) + + def hasMatch(self, query: Mol) -> Future[bool]: + """Start a query and return a future for whether any molecule matches, without waiting.""" + return self._submit(self._native.hasMatch, query) + + def getMatchesSync(self, query: Mol, maxResults: int = -1) -> list[int]: + """Like :meth:`getMatches`, but wait for and return the matching indices.""" + return self.getMatches(query, maxResults).result() + + def countMatchesSync(self, query: Mol) -> int: + """Like :meth:`countMatches`, but wait for and return the number of matches.""" + return self.countMatches(query).result() + + def hasMatchSync(self, query: Mol) -> bool: + """Like :meth:`hasMatch`, but wait for and return whether any molecule matches.""" + return self.hasMatch(query).result() diff --git a/nvmolkit/tests/test_substruct_library.py b/nvmolkit/tests/test_substruct_library.py new file mode 100644 index 00000000..64d651ed --- /dev/null +++ b/nvmolkit/tests/test_substruct_library.py @@ -0,0 +1,148 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import threading +from concurrent.futures import Future + +import pytest +from rdkit import Chem +from test_substructure import PLAIN_QUERIES, PLAIN_QUERY_TARGETS, TEST_DATA_DIR, load_smarts_file + +from nvmolkit.substruct_library import SubstructLibrary +from nvmolkit.substructure import SubstructSearchConfig + + +def test_finalize_makes_molecules_searchable(): + library = SubstructLibrary() + assert library.addMols(Chem.MolFromSmiles(smiles) for smiles in ["CCO", "c1ccccc1", "CC(=O)O"]) == [0, 1, 2] + assert len(library) == 0 + assert library.pendingSize == 3 + + query = Chem.MolFromSmarts("[#6]") + with pytest.raises(RuntimeError, match="finalized"): + library.hasMatch(query).result() + + library.finalize() + assert len(library) == 3 + assert library.pendingSize == 0 + assert library.maxConcurrentQueries >= 1 + assert library.getMatches(query).result() == [0, 1, 2] + assert library.getMatches(query, maxResults=2).result() == [0, 1] + assert library.countMatches(query).result() == 3 + assert library.hasMatch(Chem.MolFromSmarts("c")).result() + + futures = [library.hasMatch(Chem.MolFromSmarts(pattern)) for pattern in ["C", "N", "O"]] + assert all(isinstance(future, Future) for future in futures) + assert [future.result() for future in futures] == [True, False, True] + + assert library.addMol(Chem.MolFromSmiles("N")) == 3 + assert library.getMatches(query).result() == [0, 1, 2] + library.finalize() + assert len(library) == 4 + + +class _FailingFinalizeNative: + """Delegate to a native library but fail finalize, as running out of GPU memory would. + + Native rollback is covered by the C++ tests; this checks the wrapper's own executor handling. + """ + + def __init__(self, native): + self._native = native + + def __getattr__(self, name): + return getattr(self._native, name) + + def finalize(self): + raise RuntimeError("simulated finalize failure") + + +def test_wrapper_keeps_serving_queries_after_a_failed_finalize(): + library = SubstructLibrary() + library.addMols([Chem.MolFromSmiles("c1ccccc1")]) + library.finalize() + library.addMol(Chem.MolFromSmiles("Oc1ccccc1")) + + native = library._native + library._native = _FailingFinalizeNative(native) + with pytest.raises(RuntimeError, match="simulated"): + library.finalize() + library._native = native + + query = Chem.MolFromSmarts("c") + assert library.getMatches(query).result() == [0] + library.finalize() + assert library.getMatches(query).result() == [0, 1] + + +def test_future_callbacks_can_finalize_and_query(): + library = SubstructLibrary() + library.addMol(Chem.MolFromSmiles("CCO")) + library.finalize() + query = Chem.MolFromSmarts("[#8]") + finished = threading.Event() + seen = [] + + def grow_and_requery(_future): + library.addMol(Chem.MolFromSmiles("OCCO")) + library.finalize() + seen.append(library.getMatchesSync(query)) + finished.set() + + library.getMatches(query).add_done_callback(grow_and_requery) + assert finished.wait(timeout=60) + assert seen == [[0, 1]] + + +def test_invalid_library_configuration(): + config = SubstructSearchConfig(gpuIds=[0, 0]) + with pytest.raises(ValueError, match="unique"): + SubstructLibrary(config=config) + + +@pytest.mark.parametrize("query_smiles", PLAIN_QUERIES) +def test_plain_molecule_query_atoms_match_like_rdkit(query_smiles): + library = SubstructLibrary() + targets = [Chem.MolFromSmiles(smiles) for smiles in PLAIN_QUERY_TARGETS] + library.addMols(targets) + library.finalize() + + query = Chem.MolFromSmiles(query_smiles) + expected = [index for index, target in enumerate(targets) if target.HasSubstructMatch(query)] + assert library.getMatches(query).result() == expected + + +def test_chirality_is_ignored_except_in_smarts_queries(): + library = SubstructLibrary() + library.addMols([Chem.MolFromSmiles(smiles) for smiles in ["C[C@H](O)CC", "C[C@@H](O)CC", "CCC"]]) + library.finalize() + + assert library.getMatchesSync(Chem.MolFromSmiles("C[C@H](O)CC")) == [0, 1] + with pytest.raises(RuntimeError, match="chirality"): + library.getMatchesSync(Chem.MolFromSmarts("[C@H](C)(O)CC")) + + +def test_sync_methods_return_the_future_results(): + library = SubstructLibrary() + library.addMols(iter([Chem.MolFromSmiles(smiles) for smiles in ["CCO", "CCN", "OCCO"]])) + library.finalize() + + query = Chem.MolFromSmarts("[OX2H]") + assert library.getMatchesSync(query) == library.getMatches(query).result() == [0, 2] + assert library.getMatchesSync(query, maxResults=1) == [0] + assert library.countMatchesSync(query) == 2 + assert library.hasMatchSync(query) + assert not library.hasMatchSync(Chem.MolFromSmarts("[Si]")) + + +def test_matches_rdkit_on_real_molecules_and_query_sets(one_hundred_mols): + library = SubstructLibrary() + assert library.addMols(one_hundred_mols) == list(range(len(one_hundred_mols))) + library.finalize() + + queries, smarts = load_smarts_file(TEST_DATA_DIR / "SMARTS" / "rdkit_fragment_descriptors_supported.txt") + futures = [(library.getMatches(query), library.getMatches(query, maxResults=3)) for query in queries] + for query, pattern, (all_matches, limited) in zip(queries, smarts, futures): + expected = [index for index, mol in enumerate(one_hundred_mols) if mol.HasSubstructMatch(query)] + assert all_matches.result() == expected, pattern + assert limited.result() == expected[:3], pattern