Support similar and dissimilar threshold searches

This commit is contained in:
Emil
2026-07-13 23:34:58 +03:00
parent 9c3965a449
commit 134c1010f1
6 changed files with 117 additions and 11 deletions
+14 -1
View File
@@ -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 \
+11 -1
View File
@@ -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 \\
+15 -2
View File
@@ -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 = (
+50 -7
View File
@@ -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."
+8
View File
@@ -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)
+19
View File
@@ -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")
)