mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +08:00
- llm_prompt.py: RAG prompt encoding (use_rag), Qwen support, set_retrieval_pool (leak-safe, train-only pool) - models.py: use_rag/rag_top_k passthrough to LLMPromptEncoder - nested_cv_optuna.py: --use-rag CLI, per-fold retrieval pool (防泄漏), slim checkpoint (skip frozen LLM weights) - Qwen+RAG results: seed42 delivery R2=0.2151, seed123=0.1558 (mean 0.185), best LLM encoder so far
225 lines
10 KiB
Python
225 lines
10 KiB
Python
"""LLM 分子特征分支:用 MolT5 直接编码 SMILES 文本。"""
|
||
|
||
import os
|
||
from typing import Dict, List, Optional
|
||
|
||
import torch
|
||
import torch.nn as nn
|
||
|
||
|
||
# 权重默认路径,可用环境变量 MOLT5_PATH 覆盖
|
||
DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
|
||
|
||
|
||
class LLMPromptEncoder(nn.Module):
|
||
"""用 MolT5 编码 SMILES 文本,输出分子特征 F_llm [B, d_model]。
|
||
|
||
- 冻结 encoder 时:对每个 SMILES 缓存其句向量,避免重复前向,降方差、提速。
|
||
- use_lora=True 时:encoder 可训练(LoRA),不缓存。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
d_model: int,
|
||
n_chem_tokens: int = 4,
|
||
n_cond_tokens: int = 4,
|
||
model_name_or_path: str = DEFAULT_MOLT5_PATH,
|
||
freeze: bool = True,
|
||
use_lora: bool = False,
|
||
lora_r: int = 8,
|
||
lora_alpha: int = 16,
|
||
lora_dropout: float = 0.05,
|
||
max_length: int = 128,
|
||
use_rag: bool = False,
|
||
rag_top_k: int = 4,
|
||
) -> None:
|
||
super().__init__()
|
||
from transformers import AutoTokenizer, T5EncoderModel, AutoModel
|
||
|
||
self.use_lora = use_lora
|
||
self.max_length = max_length
|
||
_name_l = model_name_or_path.lower()
|
||
_is_t5 = "t5" in _name_l
|
||
_is_qwen = "qwen" in _name_l
|
||
self.use_rag = use_rag
|
||
self.rag_top_k = rag_top_k
|
||
self._is_qwen = _is_qwen
|
||
# Qwen 等大模型 tokenizer 需要 trust_remote_code
|
||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||
model_name_or_path, trust_remote_code=_is_qwen)
|
||
if _is_qwen and self.tokenizer.pad_token is None:
|
||
self.tokenizer.pad_token = self.tokenizer.eos_token
|
||
# 自动判断架构:T5→T5EncoderModel;Qwen(decoder)→AutoModel(fp16);其他→AutoModel
|
||
if _is_t5:
|
||
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
|
||
self.hidden_size = self.encoder.config.d_model
|
||
elif _is_qwen:
|
||
self.encoder = AutoModel.from_pretrained(
|
||
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16)
|
||
self.hidden_size = self.encoder.config.hidden_size
|
||
else:
|
||
self.encoder = AutoModel.from_pretrained(model_name_or_path)
|
||
self.hidden_size = self.encoder.config.hidden_size
|
||
|
||
if use_lora:
|
||
self._apply_lora(lora_r, lora_alpha, lora_dropout)
|
||
elif freeze:
|
||
for p in self.encoder.parameters():
|
||
p.requires_grad = False
|
||
self._frozen = freeze and not use_lora
|
||
|
||
self.proj_down = nn.Sequential(
|
||
nn.Linear(self.hidden_size, d_model), nn.LayerNorm(d_model)
|
||
)
|
||
# 冻结特征缓存:smiles -> [H](CPU)
|
||
self._cache: Dict[str, torch.Tensor] = {}
|
||
# RAG 检索池(防泄漏:由 nested_cv 在每个 outer fold 用训练集设置)
|
||
self._rag_pool_smiles: List[str] = []
|
||
self._rag_pool_labels = None # np.ndarray [N]
|
||
self._rag_pool_fps = None # 检索池指纹
|
||
self._rag_pool_id: str = "none" # 池子标识,用于缓存隔离
|
||
|
||
def _apply_lora(self, r: int, alpha: int, dropout: float) -> None:
|
||
from peft import LoraConfig, get_peft_model
|
||
|
||
for p in self.encoder.parameters():
|
||
p.requires_grad = False
|
||
# 不同架构注意力层命名不同:
|
||
# T5: q/k/v/o;Roberta/ChemBERTa: query/key/value;Qwen/Llama: q_proj/k_proj/v_proj/o_proj
|
||
_enc_name = type(self.encoder).__name__.lower()
|
||
if getattr(self, "_is_qwen", False) or "qwen" in _enc_name or "llama" in _enc_name:
|
||
_targets = ["q_proj", "k_proj", "v_proj", "o_proj"]
|
||
elif "t5" in _enc_name or hasattr(self.encoder.config, "d_model"):
|
||
_targets = ["q", "k", "v", "o"]
|
||
else:
|
||
_targets = ["query", "key", "value"]
|
||
cfg = LoraConfig(
|
||
r=r, lora_alpha=alpha, lora_dropout=dropout,
|
||
target_modules=_targets, bias="none",
|
||
)
|
||
self.encoder = get_peft_model(self.encoder, cfg)
|
||
|
||
def set_retrieval_pool(self, smiles_list, labels, pool_id: str = "train") -> None:
|
||
"""设置 RAG 检索池(防泄漏关键:只传训练集的 smiles 和 delivery 标签)。
|
||
labels: array-like [N] delivery 标签。pool_id: 池标识,切换时清 RAG 缓存。"""
|
||
import numpy as np
|
||
from lnp_ml.modeling.retrieval import _smiles_to_fp
|
||
self._rag_pool_smiles = list(smiles_list)
|
||
self._rag_pool_labels = np.asarray(labels, dtype=np.float32).reshape(-1)
|
||
self._rag_pool_fps = [_smiles_to_fp(s) for s in self._rag_pool_smiles]
|
||
if pool_id != self._rag_pool_id:
|
||
# 池子变了,清掉 RAG 编码缓存(避免跨 fold 串用)
|
||
self._cache = {k: v for k, v in self._cache.items() if not k.startswith("RAG::")}
|
||
self._rag_pool_id = pool_id
|
||
|
||
def _retrieve_topk(self, query_smiles: str):
|
||
"""检索 Top-K 相似邻居,返回 [(smiles, sim, label)]。排除查询分子自己(防泄漏)。"""
|
||
from rdkit import DataStructs
|
||
from lnp_ml.modeling.retrieval import _smiles_to_fp
|
||
qfp = _smiles_to_fp(query_smiles)
|
||
if qfp is None or self._rag_pool_fps is None:
|
||
return []
|
||
sims = []
|
||
for j, (smi, fp) in enumerate(zip(self._rag_pool_smiles, self._rag_pool_fps)):
|
||
if fp is None or smi == query_smiles: # 排除自己
|
||
continue
|
||
sims.append((j, DataStructs.TanimotoSimilarity(qfp, fp)))
|
||
sims.sort(key=lambda t: t[1], reverse=True)
|
||
out = []
|
||
for j, sim in sims[:self.rag_top_k]:
|
||
out.append((self._rag_pool_smiles[j], sim, float(self._rag_pool_labels[j])))
|
||
return out
|
||
|
||
def _build_rag_prompt(self, target_smiles: str, neighbors) -> str:
|
||
"""构造 RAG prompt(与探针验证版一致,已验证 R²=0.1357)。"""
|
||
blocks = []
|
||
for rank, (smi, sim, lbl) in enumerate(neighbors, 1):
|
||
blocks.append(
|
||
f"Retrieved sample {rank}:\nSMILES: {smi}\n"
|
||
f"Similarity score: {sim:.3f}\nKnown outcome - delivery_log: {lbl:.3f}"
|
||
)
|
||
retrieved_block = "\n\n".join(blocks) if blocks else "(no retrieved samples)"
|
||
return (
|
||
"Task: Encode the target LNP molecule into a retrieval-aware representation "
|
||
"for downstream delivery prediction. Do not output predictions.\n\n"
|
||
f"[Target Molecule]\nSMILES: {target_smiles}\n\n"
|
||
"[Retrieved Similar LNP Samples]\n"
|
||
"The following are retrieved from the training set by fingerprint similarity, "
|
||
"with their known delivery outcomes:\n\n"
|
||
f"{retrieved_block}\n\n"
|
||
"[Encoding Instructions]\n"
|
||
"Capture the target molecular structure and the retrieval evidence "
|
||
"(structural similarity and consistency of retrieved delivery outcomes) "
|
||
"into your internal representation."
|
||
)
|
||
|
||
@torch.no_grad()
|
||
def _encode_rag(self, smiles, device):
|
||
"""RAG 编码:每个分子→检索→构造prompt→Qwen编码→取最后有效token。带缓存。"""
|
||
feats = []
|
||
to_compute = []
|
||
keys = []
|
||
for s in smiles:
|
||
key = f"RAG::{self._rag_pool_id}::{s}"
|
||
keys.append(key)
|
||
if key not in self._cache:
|
||
to_compute.append(s)
|
||
# 逐个构造 prompt(检索是 CPU 操作)
|
||
if to_compute:
|
||
uniq = list(dict.fromkeys(to_compute))
|
||
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
|
||
for i in range(0, len(uniq), 4):
|
||
bt = prompts[i:i+4]
|
||
enc = self.tokenizer(
|
||
bt, padding=True, truncation=True,
|
||
max_length=self.max_length, return_tensors="pt",
|
||
).to(device)
|
||
out = self.encoder(**enc).last_hidden_state # [B,L,H]
|
||
lengths = enc["attention_mask"].sum(1) - 1 # 最后有效token位置
|
||
b = torch.arange(out.size(0), device=device)
|
||
last_tok = out[b, lengths.long(), :] # [B,H]
|
||
for s, v in zip(uniq[i:i+4], last_tok):
|
||
self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu()
|
||
return torch.stack([self._cache[k] for k in keys]).to(device)
|
||
|
||
def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
||
m = mask.unsqueeze(-1).float() # [B, L, 1]
|
||
return (last_hidden * m).sum(1) / m.sum(1).clamp(min=1e-6)
|
||
|
||
@torch.no_grad()
|
||
def _encode_frozen(self, smiles: List[str], device: torch.device) -> torch.Tensor:
|
||
missing = [s for s in smiles if s not in self._cache]
|
||
if missing:
|
||
uniq = list(dict.fromkeys(missing))
|
||
for i in range(0, len(uniq), 256):
|
||
chunk = uniq[i:i + 256]
|
||
enc = self.tokenizer(
|
||
chunk, padding=True, truncation=True,
|
||
max_length=self.max_length, return_tensors="pt",
|
||
).to(device)
|
||
out = self.encoder(**enc).last_hidden_state
|
||
pooled = self._mean_pool(out, enc["attention_mask"])
|
||
for s, v in zip(chunk, pooled):
|
||
self._cache[s] = v.cpu()
|
||
return torch.stack([self._cache[s] for s in smiles]).to(device)
|
||
|
||
def _encode_trainable(self, smiles: List[str], device: torch.device) -> torch.Tensor:
|
||
enc = self.tokenizer(
|
||
list(smiles), padding=True, truncation=True,
|
||
max_length=self.max_length, return_tensors="pt",
|
||
).to(device)
|
||
out = self.encoder(**enc).last_hidden_state
|
||
return self._mean_pool(out, enc["attention_mask"])
|
||
|
||
def forward(self, smiles: List[str], tab: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||
"""Args: smiles [B] SMILES 字符串列表。Returns: [B, d_model]。"""
|
||
device = self.proj_down[0].weight.device
|
||
if self.use_rag:
|
||
feat = self._encode_rag(smiles, device)
|
||
else:
|
||
feat = self._encode_frozen(smiles, device) if self._frozen \
|
||
else self._encode_trainable(smiles, device)
|
||
return self.proj_down(feat)
|
||
|
||
def clear_cache(self) -> None:
|
||
self._cache.clear() |