Compare commits

..

4 Commits

Author SHA1 Message Date
DicongLi
987ecf6386 feat: 同步 RAG 融合与训练核心代码(fusion 支持 f_retr、models/trainer/moe/dataset) 2026-07-05 01:07:26 +08:00
DicongLi
3898fb8973 feat(rag): add retrieval module (_smiles_to_fp, Top-K neighbor search) 2026-07-05 01:01:22 +08:00
Michelle0474
07b6b44f50 feat(cv): 断点续跑 - 主循环跳过已完成的 outer fold
启动时若 fold_dir 已有 test_metrics/best_params/epoch_mean,直接读回结果跳过训练。
仅对主循环生效(precomputed_best_params is None),不影响 repeat 阶段。
应对服务器偶发崩溃,崩溃重启后重跑同一 output-dir 可接续未完成的 fold。
2026-07-03 17:37:34 +08:00
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
5 changed files with 639 additions and 407 deletions

View File

@ -128,6 +128,7 @@ class ResidualConcatFusion(nn.Module):
# 零初始化门控(可学习标量),旁路初始不参与
self.g_moe = nn.Parameter(torch.zeros(()))
self.g_llm = nn.Parameter(torch.zeros(()))
self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控
def forward(
self,
@ -135,6 +136,7 @@ class ResidualConcatFusion(nn.Module):
tab: torch.Tensor,
f_moe: Optional[torch.Tensor] = None,
f_llm: Optional[torch.Tensor] = None,
f_retr: Optional[torch.Tensor] = None,
return_attn_weights: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
# 只对真实 tokenchem + tab做注意力池化旁路不参与 softmax 竞争
@ -148,5 +150,7 @@ class ResidualConcatFusion(nn.Module):
out = out + self.g_moe * f_moe # 残差 + 零初始化门
if f_llm is not None:
out = out + self.g_llm * f_llm
if f_retr is not None:
out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开
return (out, attn) if return_attn_weights else out

View File

@ -54,6 +54,17 @@ class LLMPromptEncoder(nn.Module):
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
@ -82,6 +93,10 @@ class LLMPromptEncoder(nn.Module):
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
# 不同架构注意力层命名不同:
@ -153,34 +168,39 @@ class LLMPromptEncoder(nn.Module):
"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):
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 # 最后有效token位置
lengths = enc["attention_mask"].sum(1) - 1
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):
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]

View File

@ -632,6 +632,21 @@ def _run_single_outer_fold(
fold_dir.mkdir(parents=True, exist_ok=True)
full_dataset = LNPDataset(df)
# === 断点续跑:已完成的 fold 直接跳过(读回磁盘结果)===
_tm = fold_dir / "test_metrics.json"
_bp = fold_dir / "best_params.json"
_em = fold_dir / "epoch_mean.json"
if precomputed_best_params is None and _tm.exists() and _bp.exists() and _em.exists():
logger.success(f"[SKIP] Outer fold {outer_fold} already done, loading cached results.")
with open(_tm) as _f: _tmd = json.load(_f)
with open(_bp) as _f: _bpd = json.load(_f)
with open(_em) as _f: _emd = json.load(_f)
return {
"fold": outer_fold,
"best_params": _bpd,
"epoch_mean": int(_emd.get("epoch_mean", _emd) if isinstance(_emd, dict) else _emd),
"test_metrics": _tmd,
}
logger.info(f"\n{'='*60}")
logger.info(f"OUTER FOLD {outer_fold}")

View 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]
)