lnp_ml/lnp_ml/modeling/layers/llm_prompt.py
Michelle0474 7248148515 feat(llm): Qwen QLoRA 4-bit 真微调路径 + RAG 编码冻结/微调分流
- 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 显著提升
2026-07-01 15:52:20 +08:00

245 lines
12 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:
# 微调(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/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."
)
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()