mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +08:00
feat(rag): add retrieval module (_smiles_to_fp, Top-K neighbor search)
This commit is contained in:
parent
07b6b44f50
commit
3898fb8973
193
lnp_ml/modeling/retrieval.py
Normal file
193
lnp_ml/modeling/retrieval.py
Normal file
@ -0,0 +1,193 @@
|
|||||||
|
"""
|
||||||
|
检索增强特征模块 (第一层优化)。
|
||||||
|
|
||||||
|
核心思想 (借鉴 MolRAG 的化学硬先验「结构相似 → 性质相似」):
|
||||||
|
对每个分子,用 Morgan 指纹在「训练集检索池」中找最相似的 k 个邻居,
|
||||||
|
把它们的已知标签 (已 z-score 的 quantified_delivery 等) 聚合成一个特征向量,
|
||||||
|
作为额外旁路注入融合层 (配零初始化门控,保证「加了不会更差」)。
|
||||||
|
|
||||||
|
防数据泄漏铁律 (本模块通过接口与逻辑强制保证):
|
||||||
|
1. 检索池 (pool_smiles / pool_labels) 只能由调用方传入「训练集」分子。
|
||||||
|
本模块自身不接触任何全局数据,只用传进来的池。
|
||||||
|
2. exclude_self=True 时,查询分子若在池中 (训练分子查询自己),
|
||||||
|
排除指纹完全相同的「自己」(相似度=1.0 的那一个),否则等于直接看答案。
|
||||||
|
3. 测试分子查询时,池里全是训练集,测试分子不在池中,
|
||||||
|
天然不会检索到自己或其他测试分子。
|
||||||
|
|
||||||
|
设计原则:
|
||||||
|
- 纯特征工程,不引入可训练参数 (投影层在 models.py 侧)。
|
||||||
|
- 池构建时一次性预计算所有指纹,查询时只算查询分子指纹 + 相似度,高效。
|
||||||
|
- 标签已在 dataset.py 标准化 (z-score),聚合时直接 (加权) 平均即可。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import List, Optional, Sequence
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from rdkit import Chem
|
||||||
|
from rdkit.Chem import AllChem, DataStructs
|
||||||
|
from rdkit import RDLogger
|
||||||
|
|
||||||
|
# 关闭 RDKit 的噪音日志 (无效 SMILES 警告等)
|
||||||
|
RDLogger.DisableLog("rdApp.*")
|
||||||
|
|
||||||
|
|
||||||
|
def _smiles_to_fp(smiles: str, radius: int = 2, n_bits: int = 2048):
|
||||||
|
"""单个 SMILES -> Morgan 指纹 (ExplicitBitVect)。无效 SMILES 返回 None。"""
|
||||||
|
mol = Chem.MolFromSmiles(smiles)
|
||||||
|
if mol is None:
|
||||||
|
return None
|
||||||
|
return AllChem.GetMorganFingerprintAsBitVect(mol, radius=radius, nBits=n_bits)
|
||||||
|
|
||||||
|
|
||||||
|
class MorganRetriever:
|
||||||
|
"""基于 Morgan 指纹 + Tanimoto 相似度的分子检索器。
|
||||||
|
|
||||||
|
用法:
|
||||||
|
# 只用训练集分子建池
|
||||||
|
retr = MorganRetriever(
|
||||||
|
pool_smiles=train_smiles, # List[str]
|
||||||
|
pool_labels=train_labels, # np.ndarray [N_pool, label_dim]
|
||||||
|
radius=2, n_bits=2048, k=5,
|
||||||
|
)
|
||||||
|
# 查询 (训练分子查询时 exclude_self=True 防泄漏)
|
||||||
|
feat = retr.query(query_smiles, exclude_self=True) # np.ndarray [label_dim*?]
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
pool_smiles: Sequence[str],
|
||||||
|
pool_labels: np.ndarray,
|
||||||
|
radius: int = 2,
|
||||||
|
n_bits: int = 2048,
|
||||||
|
k: int = 5,
|
||||||
|
sim_weighted: bool = True,
|
||||||
|
include_sim_stats: bool = True,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
pool_smiles: 检索池分子的 SMILES (必须只含训练集分子)。
|
||||||
|
pool_labels: 检索池分子的标签, shape [N_pool, label_dim], 已标准化。
|
||||||
|
radius: Morgan 指纹半径 (MolRAG 用 2)。
|
||||||
|
n_bits: 指纹位数。
|
||||||
|
k: 检索近邻数。
|
||||||
|
sim_weighted: 聚合邻居标签时是否按相似度加权 (否则等权平均)。
|
||||||
|
include_sim_stats: 是否在输出特征里附带相似度统计 (top1/mean 相似度),
|
||||||
|
让模型知道「邻居有多可靠」。
|
||||||
|
"""
|
||||||
|
assert len(pool_smiles) == len(pool_labels), "池 SMILES 与标签数量不一致"
|
||||||
|
self.radius = radius
|
||||||
|
self.n_bits = n_bits
|
||||||
|
self.k = k
|
||||||
|
self.sim_weighted = sim_weighted
|
||||||
|
self.include_sim_stats = include_sim_stats
|
||||||
|
|
||||||
|
_smiles_arr = list(pool_smiles)
|
||||||
|
_labels_arr = np.asarray(pool_labels, dtype=np.float32)
|
||||||
|
# 过滤掉标签含 NaN 的分子:它们没有有效标签,不能作为性质参考邻居
|
||||||
|
_valid_mask = ~np.isnan(_labels_arr).any(axis=1)
|
||||||
|
self.pool_smiles: List[str] = [_smiles_arr[i] for i in range(len(_smiles_arr)) if _valid_mask[i]]
|
||||||
|
self.pool_labels = _labels_arr[_valid_mask]
|
||||||
|
self.label_dim = self.pool_labels.shape[1]
|
||||||
|
_n_dropped = int((~_valid_mask).sum())
|
||||||
|
if _n_dropped > 0:
|
||||||
|
import sys
|
||||||
|
print(f'[MorganRetriever] 过滤 {_n_dropped} 个 NaN 标签分子,有效检索池={len(self.pool_smiles)}', file=sys.stderr)
|
||||||
|
|
||||||
|
# 池内 SMILES 集合,用于 exclude_self="auto" 时判断查询分子是否在池中
|
||||||
|
self.pool_smiles_set = set(self.pool_smiles)
|
||||||
|
|
||||||
|
# 预计算池内所有指纹 (无效分子记 None,检索时跳过)
|
||||||
|
self.pool_fps = [_smiles_to_fp(s, radius, n_bits) for s in self.pool_smiles]
|
||||||
|
|
||||||
|
# 标签均值,用于无有效邻居时的回退 (用训练集均值,不泄漏)
|
||||||
|
self.label_mean = self.pool_labels.mean(axis=0)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def feature_dim(self) -> int:
|
||||||
|
"""输出特征维度: 聚合标签 (label_dim) [+ 相似度统计 2 维]。"""
|
||||||
|
return self.label_dim + (2 if self.include_sim_stats else 0)
|
||||||
|
|
||||||
|
def query(self, query_smiles: str, exclude_self="auto") -> np.ndarray:
|
||||||
|
"""检索查询分子的 top-k 邻居并聚合其标签为特征向量。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query_smiles: 查询分子 SMILES。
|
||||||
|
exclude_self: 排除自身策略,防数据泄漏的关键。
|
||||||
|
- "auto" (默认, 推荐): 查询分子若在检索池中(训练分子),自动排除自己;
|
||||||
|
不在池中(测试分子)则不排除。无需调用方区分训练/测试。
|
||||||
|
- True: 强制排除相似度=1.0 的分子。
|
||||||
|
- False: 不排除(仅当确定查询分子不在池中时使用)。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
特征向量 np.ndarray [feature_dim]。无有效邻居时回退到训练集均值。
|
||||||
|
"""
|
||||||
|
# auto 模式:查询分子在池中 → 训练分子 → 排除自己;否则不排除
|
||||||
|
if exclude_self == "auto":
|
||||||
|
do_exclude = query_smiles in self.pool_smiles_set
|
||||||
|
else:
|
||||||
|
do_exclude = bool(exclude_self)
|
||||||
|
|
||||||
|
q_fp = _smiles_to_fp(query_smiles, self.radius, self.n_bits)
|
||||||
|
if q_fp is None:
|
||||||
|
# 查询分子无效:回退到训练集均值 + 零相似度
|
||||||
|
agg = self.label_mean.copy()
|
||||||
|
if self.include_sim_stats:
|
||||||
|
agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)])
|
||||||
|
return agg.astype(np.float32)
|
||||||
|
|
||||||
|
# 与池内每个分子算 Tanimoto 相似度
|
||||||
|
sims = np.full(len(self.pool_fps), -1.0, dtype=np.float32)
|
||||||
|
for i, fp in enumerate(self.pool_fps):
|
||||||
|
if fp is None:
|
||||||
|
continue
|
||||||
|
sims[i] = DataStructs.TanimotoSimilarity(q_fp, fp)
|
||||||
|
|
||||||
|
# 排除自己: 指纹完全相同 (相似度 >= 1.0 - eps) 的那一个
|
||||||
|
if do_exclude:
|
||||||
|
eps = 1e-6
|
||||||
|
self_mask = sims >= (1.0 - eps)
|
||||||
|
# 保守起见:把所有相似度=1.0 的都视作潜在自身并排除,
|
||||||
|
# 因为结构完全相同的分子标签也应相同,留着等于泄漏答案。
|
||||||
|
sims[self_mask] = -1.0
|
||||||
|
|
||||||
|
# 取 top-k (相似度降序),过滤掉无效 (-1) 的
|
||||||
|
valid = np.where(sims >= 0.0)[0]
|
||||||
|
if len(valid) == 0:
|
||||||
|
# 没有有效邻居:回退到训练集均值
|
||||||
|
agg = self.label_mean.copy()
|
||||||
|
if self.include_sim_stats:
|
||||||
|
agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)])
|
||||||
|
return agg.astype(np.float32)
|
||||||
|
|
||||||
|
order = valid[np.argsort(-sims[valid])]
|
||||||
|
topk_idx = order[: self.k]
|
||||||
|
topk_sims = sims[topk_idx]
|
||||||
|
topk_labels = self.pool_labels[topk_idx] # [k', label_dim]
|
||||||
|
|
||||||
|
# 聚合邻居标签
|
||||||
|
if self.sim_weighted and topk_sims.sum() > 1e-8:
|
||||||
|
w = topk_sims / topk_sims.sum()
|
||||||
|
agg_label = (w[:, None] * topk_labels).sum(axis=0)
|
||||||
|
else:
|
||||||
|
agg_label = topk_labels.mean(axis=0)
|
||||||
|
|
||||||
|
if self.include_sim_stats:
|
||||||
|
sim_stats = np.array(
|
||||||
|
[float(topk_sims[0]), float(topk_sims.mean())], dtype=np.float32
|
||||||
|
)
|
||||||
|
agg = np.concatenate([agg_label.astype(np.float32), sim_stats])
|
||||||
|
else:
|
||||||
|
agg = agg_label.astype(np.float32)
|
||||||
|
|
||||||
|
return agg.astype(np.float32)
|
||||||
|
|
||||||
|
def query_batch(
|
||||||
|
self, query_smiles_list, exclude_self="auto"
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""批量检索。返回 [N_query, feature_dim]。"""
|
||||||
|
return np.stack(
|
||||||
|
[self.query(s, exclude_self=exclude_self) for s in query_smiles_list]
|
||||||
|
)
|
||||||
Loading…
x
Reference in New Issue
Block a user