From 134c1010f1b3092c3746208d3e296e0329c6798f Mon Sep 17 00:00:00 2001 From: Emil Date: Mon, 13 Jul 2026 23:34:58 +0300 Subject: [PATCH] Support similar and dissimilar threshold searches --- README.md | 15 ++++++- scimesh/workloads/help.py | 12 +++++- scimesh/workloads/similarity_graph.py | 17 +++++++- scimesh/workloads/similarity_search.py | 57 ++++++++++++++++++++++---- tests/test_similarity_graph.py | 8 ++++ tests/test_similarity_search.py | 19 +++++++++ 6 files changed, 117 insertions(+), 11 deletions(-) diff --git a/README.md b/README.md index f405a32..9ddef99 100644 --- a/README.md +++ b/README.md @@ -60,6 +60,19 @@ scimesh similarity-search chembl_37_chemreps.txt \ The output CSV contains `rank,chembl_id,canonical_smiles,similarity`. Search progress and valid/invalid-SMILES statistics are written to the terminal. `--max-rows` limits the candidate scan for small tests, while `--progress-every 0` disables progress reports. +To find the least similar molecules, use `--threshold-direction less`. This ranks +results from the lowest similarity upward; `--threshold` optionally limits them +to values less than or equal to a cutoff: + +```bash +scimesh similarity-search chembl_37_chemreps.txt \ + --query-id CHEMBL939 \ + --threshold-direction less \ + --threshold 0.1 \ + --top-k 20 \ + --output least_similar.csv +``` + To render the query and retained candidates: ```bash @@ -72,7 +85,7 @@ This writes `query.png` and `top_candidates.png` into `structures`. ## Similarity graph -`similarity-graph` constructs an exact sparse undirected graph. Every valid molecule is a vertex; an edge is emitted only when Tanimoto similarity is at least `--threshold`. Each fingerprint is calculated once. Comparisons are processed block by block, each pair is tested once (`i < j`), and no dense N×N matrix is created or stored. +`similarity-graph` constructs an exact sparse undirected graph. Every valid molecule is a vertex; an edge is emitted when Tanimoto similarity satisfies the selected threshold direction (`>=` by default, or `<=` with `--threshold-direction less`). Each fingerprint is calculated once. Comparisons are processed block by block, each pair is tested once (`i < j`), and no dense N×N matrix is created or stored. ```bash scimesh similarity-graph chembl_37_chemreps.txt \ diff --git a/scimesh/workloads/help.py b/scimesh/workloads/help.py index a220195..d361fad 100644 --- a/scimesh/workloads/help.py +++ b/scimesh/workloads/help.py @@ -34,7 +34,17 @@ HELP_TEXT = dedent( --top-k 20 \\ --output results/smiles_search.csv - 4. Build a small exact similarity graph: + 4. Find the least similar molecules. With "less", results are ranked from + lowest similarity upward; --threshold is an optional <= filter: + + scimesh similarity-search chembl_37_chemreps.txt \\ + --query-id CHEMBL939 \\ + --threshold-direction less \\ + --threshold 0.1 \\ + --top-k 20 \\ + --output results/least_similar.csv + + 5. Build a small exact similarity graph: scimesh similarity-graph chembl_37_chemreps.txt \\ --max-rows 1000 \\ diff --git a/scimesh/workloads/similarity_graph.py b/scimesh/workloads/similarity_graph.py index bc84d7c..4c6fa0d 100644 --- a/scimesh/workloads/similarity_graph.py +++ b/scimesh/workloads/similarity_graph.py @@ -60,12 +60,15 @@ def build_similarity_graph( block_size: int, max_rows: int | None = None, progress_every: int = 0, + threshold_direction: str = "greater", ) -> 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") + if threshold_direction not in {"greater", "less"}: + raise ValueError("--threshold-direction must be 'greater' or 'less'") molecules, stats = _fingerprinted_molecules(tsv_path, max_rows) edges: list[SimilarityEdge] = [] @@ -86,7 +89,12 @@ def build_similarity_graph( molecules[left_index].fingerprint, molecules[right_index].fingerprint, ) - if similarity >= threshold: + matches_threshold = ( + similarity >= threshold + if threshold_direction == "greater" + else similarity <= threshold + ) + if matches_threshold: edges.append( SimilarityEdge( molecules[left_index].molecule_id, @@ -134,7 +142,11 @@ class SimilarityGraphWorkload: 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", + help="Similarity value used to create graph edges", + ) + parser.add_argument( + "--threshold-direction", choices=("greater", "less"), default="greater", + help="Create edges with >= threshold or <= threshold similarity", ) parser.add_argument( "--block-size", type=int, default=1_000, @@ -164,6 +176,7 @@ class SimilarityGraphWorkload: args.block_size, args.max_rows, args.progress_every, + args.threshold_direction, ) write_graph_edges(args.output, result.edges) rate = ( diff --git a/scimesh/workloads/similarity_search.py b/scimesh/workloads/similarity_search.py index 7d1167d..d3c92de 100644 --- a/scimesh/workloads/similarity_search.py +++ b/scimesh/workloads/similarity_search.py @@ -31,7 +31,10 @@ class SimilarityMatch: molecule_id: str smiles: str - def sort_key(self) -> tuple[float, str, str]: + def sort_key(self, threshold_direction: str = "greater") -> tuple[float, str, str]: + """Return a stable ranking key for similar or dissimilar searches.""" + if threshold_direction == "less": + return (self.similarity, self.molecule_id, self.smiles) return (-self.similarity, self.molecule_id, self.smiles) @@ -40,11 +43,12 @@ class _HeapEntry: """Heap item whose minimum is the worst retained match.""" match: SimilarityMatch + rank_key: tuple[float, str, str] def __lt__(self, other: object) -> bool: if not isinstance(other, _HeapEntry): return NotImplemented - return self.match.sort_key() > other.match.sort_key() + return self.rank_key > other.rank_key @dataclass @@ -61,10 +65,16 @@ def search_similar( top_k: int, max_rows: int | None = None, progress_every: int = 0, + threshold: float | None = None, + threshold_direction: str = "greater", ) -> 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") + if threshold is not None and not 0.0 <= threshold <= 1.0: + raise ValueError("--threshold must be between 0 and 1") + if threshold_direction not in {"greater", "less"}: + raise ValueError("--threshold-direction must be 'greater' or 'less'") query_fingerprint = fingerprint(query.molecule) query_canonical_smiles = Chem.MolToSmiles(query.molecule, canonical=True) stats = DatasetStats() @@ -80,15 +90,23 @@ def search_similar( or candidate_canonical_smiles == query_canonical_smiles ): continue + similarity = DataStructs.TanimotoSimilarity( + query_fingerprint, fingerprint(record.molecule) + ) + if threshold is not None and ( + similarity < threshold if threshold_direction == "greater" else similarity > threshold + ): + continue match = SimilarityMatch( - DataStructs.TanimotoSimilarity(query_fingerprint, fingerprint(record.molecule)), + similarity, record.molecule_id, record.smiles, ) - entry = _HeapEntry(match) + rank_key = match.sort_key(threshold_direction) + entry = _HeapEntry(match, rank_key) if len(heap) < top_k: heapq.heappush(heap, entry) - elif match.sort_key() < heap[0].match.sort_key(): + elif rank_key < heap[0].rank_key: heapq.heapreplace(heap, entry) if progress_every and stats.scanned % progress_every == 0: @@ -106,7 +124,13 @@ def search_similar( last_report_at = now last_report_rows = stats.scanned - return SearchResult(sorted((entry.match for entry in heap), key=SimilarityMatch.sort_key), stats) + return SearchResult( + sorted( + (entry.match for entry in heap), + key=lambda match: match.sort_key(threshold_direction), + ), + stats, + ) def write_search_results(output_path: Path, matches: list[SimilarityMatch]) -> None: @@ -165,6 +189,14 @@ class SimilaritySearchWorkload: "--top-k", "--top", dest="top_k", type=int, default=20, help="Number of matches to retain (default: 20)", ) + parser.add_argument( + "--threshold", type=float, + help="Keep similarities at or beyond this value", + ) + parser.add_argument( + "--threshold-direction", choices=("greater", "less"), default="greater", + help="Use >= for similar molecules or <= for dissimilar molecules", + ) parser.add_argument( "-o", "--output", type=Path, default=Path("similarity_results.csv"), help="Output CSV path", @@ -202,7 +234,13 @@ class SimilaritySearchWorkload: query = MoleculeRecord("query", args.query_smiles, molecule) result = search_similar( - args.input, query, args.top_k, args.max_rows, args.progress_every + args.input, + query, + args.top_k, + args.max_rows, + args.progress_every, + args.threshold, + args.threshold_direction, ) write_search_results(args.output, result.matches) image_paths: tuple[Path, Path] | None = None @@ -212,6 +250,11 @@ class SimilaritySearchWorkload: ) print(f"Query {query.molecule_id}: {query.smiles}") + if args.threshold is not None: + operator = ">=" if args.threshold_direction == "greater" else "<=" + print(f"Similarity filter: {operator} {args.threshold:.6f}") + elif args.threshold_direction == "less": + print("Ranking least similar molecules first.") print( f"Scanned {result.stats.scanned:,} rows: {result.stats.valid:,} valid, " f"{result.stats.invalid:,} invalid SMILES." diff --git a/tests/test_similarity_graph.py b/tests/test_similarity_graph.py index 23e685b..9341e01 100644 --- a/tests/test_similarity_graph.py +++ b/tests/test_similarity_graph.py @@ -53,3 +53,11 @@ def test_graph_is_block_size_independent_and_deterministic( write_graph_edges(first_path, first.edges) write_graph_edges(second_path, repeated.edges) assert first_path.read_bytes() == second_path.read_bytes() + + +def test_graph_supports_less_than_threshold_direction(small_dataset: Path) -> None: + result = build_similarity_graph( + small_dataset, threshold=0.15, block_size=2, threshold_direction="less" + ) + + assert all(edge.similarity <= 0.15 for edge in result.edges) diff --git a/tests/test_similarity_search.py b/tests/test_similarity_search.py index a8b1538..4af1c62 100644 --- a/tests/test_similarity_search.py +++ b/tests/test_similarity_search.py @@ -37,3 +37,22 @@ def test_search_matches_full_sorting_and_skips_query_and_invalid( assert "DUPLICATE" not in {match.molecule_id for match in result.matches} assert result.stats.invalid == 1 assert result.stats.valid == 5 + + +def test_search_can_rank_and_filter_least_similar_molecules( + small_dataset: Path, +) -> None: + query = find_molecule_by_id(small_dataset, "QUERY") + result = search_similar( + small_dataset, + query, + top_k=20, + threshold=0.2, + threshold_direction="less", + ) + + assert result.matches + assert all(match.similarity <= 0.2 for match in result.matches) + assert result.matches == sorted( + result.matches, key=lambda match: match.sort_key("less") + )