mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-21 21:22:05 +08:00
85 lines
3.4 KiB
Python
85 lines
3.4 KiB
Python
"""Morgan 指纹检索:RAG prompt 构造与检索旁路共用。
|
||
|
||
- _smiles_to_fp: SMILES -> RDKit ExplicitBitVect(供 Tanimoto 相似度)。
|
||
- MorganRetriever: 训练集指纹检索器,query_batch 返回每个查询的 3 维检索特征。
|
||
"""
|
||
|
||
from typing import List, Optional
|
||
|
||
import numpy as np
|
||
from rdkit import Chem, DataStructs
|
||
from rdkit.Chem import AllChem
|
||
|
||
# 与 RDKitFeaturizer 的 morgan token 保持一致
|
||
MORGAN_RADIUS = 2
|
||
MORGAN_NBITS = 1024
|
||
|
||
|
||
def _smiles_to_fp(smiles: Optional[str], radius: int = MORGAN_RADIUS, n_bits: int = MORGAN_NBITS):
|
||
"""SMILES -> Morgan ExplicitBitVect;非法/空返回 None。"""
|
||
if not smiles:
|
||
return None
|
||
mol = Chem.MolFromSmiles(smiles)
|
||
if mol is None:
|
||
return None
|
||
return AllChem.GetMorganFingerprintAsBitVect(mol, radius=radius, nBits=n_bits)
|
||
|
||
|
||
class MorganRetriever:
|
||
"""基于 Morgan + Tanimoto 的训练集检索器(用于 use_retrieval 旁路)。
|
||
|
||
query_batch 为每个查询返回 3 维特征:
|
||
[相似度加权邻居标签均值, top1 相似度, top-k 平均相似度]
|
||
"""
|
||
|
||
def __init__(self, smiles_list: List[str], labels, k: int = 5) -> None:
|
||
self.k = int(k)
|
||
self.pool_smiles = list(smiles_list)
|
||
|
||
labels = np.asarray(labels, dtype=np.float32)
|
||
if labels.ndim == 1:
|
||
labels = labels.reshape(-1, 1)
|
||
self.pool_labels = labels # [N, D],旁路只用第 0 列
|
||
|
||
fps = [_smiles_to_fp(s) for s in self.pool_smiles]
|
||
# 只保留可解析的池分子
|
||
self._valid_pos = [i for i, fp in enumerate(fps) if fp is not None]
|
||
self._valid_fps = [fps[i] for i in self._valid_pos]
|
||
self._valid_labels = self.pool_labels[self._valid_pos, 0] if self._valid_pos else np.zeros(0, np.float32)
|
||
self._valid_smiles = [self.pool_smiles[i] for i in self._valid_pos]
|
||
self._fp_cache = {} # query 指纹缓存
|
||
|
||
def _query_fp(self, smiles: str):
|
||
if smiles not in self._fp_cache:
|
||
self._fp_cache[smiles] = _smiles_to_fp(smiles)
|
||
return self._fp_cache[smiles]
|
||
|
||
def query_batch(self, smiles_list: List[str], exclude_self="auto") -> np.ndarray:
|
||
return np.stack([self._query_one(s, exclude_self) for s in smiles_list]).astype(np.float32)
|
||
|
||
def _query_one(self, smiles: str, exclude_self) -> np.ndarray:
|
||
qfp = self._query_fp(smiles)
|
||
if qfp is None or not self._valid_fps:
|
||
return np.zeros(3, dtype=np.float32)
|
||
|
||
# 向量化 Tanimoto(C 层批量),替代 Python 逐条
|
||
sims = np.asarray(DataStructs.BulkTanimotoSimilarity(qfp, self._valid_fps), dtype=np.float32)
|
||
|
||
if exclude_self:
|
||
self_mask = np.fromiter((s == smiles for s in self._valid_smiles), dtype=bool, count=len(sims))
|
||
sims = np.where(self_mask, -1.0, sims)
|
||
|
||
k = min(self.k, int((sims >= 0).sum()))
|
||
if k <= 0:
|
||
return np.zeros(3, dtype=np.float32)
|
||
|
||
# argpartition 取 top-k(O(N)),再对这 k 个排序
|
||
top = np.argpartition(-sims, k - 1)[:k]
|
||
top = top[np.argsort(-sims[top])]
|
||
top_sims = sims[top]
|
||
top_labels = self._valid_labels[top]
|
||
|
||
w = np.clip(top_sims, 0.0, None)
|
||
wsum = float(w.sum())
|
||
weighted_label = float((w * top_labels).sum() / wsum) if wsum > 0 else float(top_labels.mean())
|
||
return np.array([weighted_label, float(top_sims[0]), float(top_sims.mean())], dtype=np.float32) |