mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-21 21:22:05 +08:00
- use_lora 时 Qwen 改 4-bit(NF4) 加载,配 prepare_model_for_kbit_training 使 LoRA 可回传梯度 - _encode_rag 拆分:冻结走 no_grad+缓存,微调走保留计算图+不缓存 - fold_0/1 delivery R2 0.44/0.24,较冻结版 0.185 显著提升
245 lines
12 KiB
Python
245 lines
12 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:
|
||
# 微调(use_lora)时用 4-bit 量化(QLoRA)省显存;纯冻结时用 fp16
|
||
if use_lora:
|
||
from transformers import BitsAndBytesConfig
|
||
_bnb = BitsAndBytesConfig(
|
||
load_in_4bit=True, bnb_4bit_quant_type="nf4",
|
||
bnb_4bit_compute_dtype=torch.float16,
|
||
bnb_4bit_use_double_quant=True)
|
||
self.encoder = AutoModel.from_pretrained(
|
||
model_name_or_path, trust_remote_code=True,
|
||
quantization_config=_bnb, device_map={"": 0})
|
||
else:
|
||
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
|
||
|
||
# 4-bit 量化模型(QLoRA)需先 prepare,才能正确接收梯度
|
||
if getattr(self.encoder, "is_loaded_in_4bit", False) or getattr(self.encoder, "is_loaded_in_8bit", False):
|
||
from peft import prepare_model_for_kbit_training
|
||
self.encoder = prepare_model_for_kbit_training(self.encoder)
|
||
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."
|
||
)
|
||
|
||
def _rag_encode_batch(self, prompts, device):
|
||
"""编码一批 RAG prompt,取最后有效 token。grad 由调用方上下文决定。"""
|
||
outs = []
|
||
for i in range(0, len(prompts), 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
|
||
b = torch.arange(out.size(0), device=device)
|
||
outs.append(out[b, lengths.long(), :].float()) # [B,H]
|
||
return torch.cat(outs, 0)
|
||
|
||
def _encode_rag(self, smiles, device):
|
||
"""RAG 编码。冻结时:no_grad + 缓存(快)。微调时:算梯度 + 不缓存。"""
|
||
if self._frozen:
|
||
# 冻结路径:缓存复用,no_grad
|
||
keys = [f"RAG::{self._rag_pool_id}::{s}" for s in smiles]
|
||
to_compute = [s for s, k in zip(smiles, keys) if k not in self._cache]
|
||
if to_compute:
|
||
uniq = list(dict.fromkeys(to_compute))
|
||
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
|
||
with torch.no_grad():
|
||
feat = self._rag_encode_batch(prompts, device)
|
||
for s, v in zip(uniq, feat):
|
||
self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu()
|
||
return torch.stack([self._cache[k] for k in keys]).to(device)
|
||
else:
|
||
# 微调路径:每次重新编码,保留计算图(算梯度),不缓存
|
||
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in smiles]
|
||
return self._rag_encode_batch(prompts, 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() |