Support similar and dissimilar threshold searches
This commit is contained in:
@@ -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 \
|
||||
|
||||
@@ -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 \\
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user