lnp_ml/lnp_ml/modeling/retrieval.py

85 lines
3.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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)
# 向量化 TanimotoC 层批量),替代 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-kO(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)