lnp_ml/lnp_ml/modeling/layers/llm_prompt.py
Michelle0474 0fc72a9474 Add Qwen2.5-7B MolRAG encoder branch + dual-seed results
- 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
2026-06-28 20:42:27 +08:00

225 lines
10 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.

"""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→T5EncoderModelQwen(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/oRoberta/ChemBERTa: query/key/valueQwen/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()