feat: integrate MolT5 LLM encoder and fix desc dim

This commit is contained in:
Michelle0474 2026-06-22 19:39:24 +08:00
parent cd4534f05a
commit 54c734dc16
4 changed files with 1279 additions and 1099 deletions

View File

@ -0,0 +1,104 @@
"""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,
) -> None:
super().__init__()
from transformers import AutoTokenizer, T5EncoderModel
self.use_lora = use_lora
self.max_length = max_length
self.tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
self.hidden_size = self.encoder.config.d_model
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] = {}
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
cfg = LoraConfig(
r=r, lora_alpha=alpha, lora_dropout=dropout,
target_modules=["q", "k", "v", "o"], bias="none",
)
self.encoder = get_peft_model(self.encoder, cfg)
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
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()

View File

@ -5,20 +5,27 @@ import torch.nn as nn
from typing import Dict, List, Optional, Literal
from lnp_ml.modeling.encoders import CachedRDKitEncoder, CachedMPNNEncoder
from lnp_ml.modeling.layers import TokenProjector, CrossModalAttention, FusionLayer, MoEBlock
from lnp_ml.modeling.layers import (
TokenProjector,
SetTransformer,
ResidualConcatFusion,
MoEBlock,
LLMPromptEncoder,
)
from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH
from lnp_ml.modeling.heads import MultiTaskHead
PoolingStrategy = Literal["concat", "avg", "max", "attention"]
PoolingStrategy = Literal["attention", "avg", "max"]
# Token 维度配置(根据 ARCHITECTURE.md
# Token 维度配置
DEFAULT_INPUT_DIMS = {
# Channel A: 化学特征
"mpnn": 600, # D-MPNN embedding
"morgan": 1024, # Morgan fingerprint
"maccs": 167, # MACCS keys
"desc": 210, # RDKit descriptors
"desc": 217, # RDKit descriptors
# Channel B: 配方/实验条件
"comp": 5, # 配方比例
"phys": 12, # 物理参数 one-hot
@ -26,8 +33,21 @@ DEFAULT_INPUT_DIMS = {
"exp": 32, # 实验条件 one-hot
}
# Token 顺序(前 4 个为 Channel A后 4 个为 Channel B
TOKEN_ORDER = ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
# 化学 / 配方 token 的键顺序
CHEM_KEYS_WITH_MPNN = ["mpnn", "morgan", "maccs", "desc"]
CHEM_KEYS_NO_MPNN = ["morgan", "maccs", "desc"]
TAB_KEYS = ["comp", "phys", "help", "exp"]
# backbone 权重前缀(用于预训练加载与导出)
BACKBONE_PREFIXES = (
"token_projector.",
"set_transformer.",
"fusion.",
"moe.",
"llm_prompt.",
)
# 冻结的 MolT5 encoder 权重前缀,不纳入 backbone由本地权重加载不进 checkpoint
LLM_FROZEN_PREFIX = "llm_prompt.encoder."
class LNPModel(nn.Module):
@ -37,37 +57,47 @@ class LNPModel(nn.Module):
架构流程:
1. Encoders: SMILES -> 化学特征; tabular -> 配方/实验特征
2. TokenProjector: 统一到 d_model
3. Stack: [B, 8, d_model]
4. CrossModalAttention: Channel A (化学) <-> Channel B (配方/实验)
5. FusionLayer: [B, 8, d_model] -> [B, fusion_dim]
6. MultiTaskHead: 多任务预测
3. SetTransformer: 对化学 token 集合做置换等变编码 -> chem'
4. MoE (可选): router tabexpert chem' -> F_moe
5. LLM (可选): chem'/tab 注入 MolT5 prompt -> F_llm
6. ResidualConcatFusion: 拼接 chem'/tab/F_moe/F_llm -> attention pooling
7. MultiTaskHead: 多任务预测
"""
def __init__(
self,
# 模型维度
d_model: int = 256,
# Cross attention
# Set Transformer
num_heads: int = 8,
n_attn_layers: int = 4,
set_transformer_block: str = "sab",
# Fusion
fusion_strategy: PoolingStrategy = "attention",
# Head
head_hidden_dim: int = 128,
# Dropout
dropout: float = 0.1,
# MPNN encoder (可选,如果不用 MPNN 可以设为 None)
# MPNN encoder
mpnn_checkpoint: Optional[str] = None,
mpnn_ensemble_paths: Optional[List[str]] = None,
mpnn_device: str = "cpu",
# 输入维度配置
input_dims: Optional[Dict[str, int]] = None,
# ============ MoE 相关(新增) ============
# ============ MoE 相关 ============
use_moe: bool = False,
moe_n_experts: int = 4,
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
# ============ LLM 相关 ============
use_llm: bool = False,
llm_model_path: str = DEFAULT_MOLT5_PATH,
llm_freeze: bool = True,
llm_use_lora: bool = False,
llm_lora_r: int = 8,
llm_lora_alpha: int = 16,
llm_lora_dropout: float = 0.05,
) -> None:
super().__init__()
@ -76,10 +106,7 @@ class LNPModel(nn.Module):
self.use_mpnn = mpnn_checkpoint is not None or mpnn_ensemble_paths is not None
# ============ Encoders ============
# RDKit encoder (always used)
self.rdkit_encoder = CachedRDKitEncoder()
# MPNN encoder (optional)
if self.use_mpnn:
self.mpnn_encoder = CachedMPNNEncoder(
checkpoint_path=mpnn_checkpoint,
@ -90,27 +117,28 @@ class LNPModel(nn.Module):
self.mpnn_encoder = None
# ============ Token Projector ============
# 根据是否使用 MPNN 调整输入维度
proj_input_dims = {k: v for k, v in self.input_dims.items()}
if not self.use_mpnn:
proj_input_dims.pop("mpnn", None)
self.token_projector = TokenProjector(
input_dims=proj_input_dims,
d_model=d_model,
dropout=dropout,
)
# ============ Cross Modal Attention ============
n_tokens = 8 if self.use_mpnn else 7
split_idx = 4 if self.use_mpnn else 3 # Channel A 的 token 数量
# token 顺序与化学侧 token 数
self.chem_keys = CHEM_KEYS_WITH_MPNN if self.use_mpnn else CHEM_KEYS_NO_MPNN
self.tab_keys = TAB_KEYS
self.token_order = self.chem_keys + self.tab_keys
self.split_idx = len(self.chem_keys)
self.cross_attention = CrossModalAttention(
# ============ Set Transformer ============
self.set_transformer = SetTransformer(
d_model=d_model,
num_heads=num_heads,
n_layers=n_attn_layers,
split_idx=split_idx,
dropout=dropout,
block=set_transformer_block,
)
# ============ MoE Block (可选) ============
@ -119,24 +147,35 @@ class LNPModel(nn.Module):
if use_moe:
self.moe = MoEBlock(
d_model=d_model,
n_chem_tokens=split_idx,
n_chem_tokens=self.split_idx,
n_experts=moe_n_experts,
top_k=moe_top_k,
expert_hidden_mult=moe_expert_hidden_mult,
dropout=dropout,
jitter_noise=moe_jitter_noise,
)
n_fusion_tokens = n_tokens + 1 # 多一个 F_moe token
else:
self.moe = None
n_fusion_tokens = n_tokens
# ============ Fusion Layer ============
self.fusion = FusionLayer(
# ============ LLM Prompt (可选) ============
self.use_llm = use_llm
if use_llm:
self.llm_prompt = LLMPromptEncoder(
d_model=d_model,
n_tokens=n_fusion_tokens,
strategy=fusion_strategy,
n_chem_tokens=self.split_idx,
n_cond_tokens=len(self.tab_keys),
model_name_or_path=llm_model_path,
freeze=llm_freeze,
use_lora=llm_use_lora,
lora_r=llm_lora_r,
lora_alpha=llm_lora_alpha,
lora_dropout=llm_lora_dropout,
)
else:
self.llm_prompt = None
# ============ Residual Concat + Fusion ============
self.fusion = ResidualConcatFusion(d_model=d_model, strategy=fusion_strategy)
# ============ Multi-Task Head ============
self.head = MultiTaskHead(
@ -151,70 +190,51 @@ class LNPModel(nn.Module):
tabular: Dict[str, torch.Tensor],
) -> torch.Tensor:
"""
内部方法编码 SMILES tabular返回 stacked tokens
编码 SMILES tabular返回 stacked tokens
Returns:
stacked: [B, n_tokens, d_model]
stacked: [B, n_tokens, d_model]顺序为 chem 在前tab 在后
"""
# 获取目标设备(从 tabular 数据推断)
device = tabular["comp"].device
# 1. Encode SMILES
rdkit_features = self.rdkit_encoder(smiles) # {"morgan", "maccs", "desc"}
rdkit_features = self.rdkit_encoder(smiles)
# 2. 合并所有特征
all_features: Dict[str, torch.Tensor] = {}
# MPNN 特征(如果启用)
if self.use_mpnn:
mpnn_features = self.mpnn_encoder(smiles)
all_features["mpnn"] = mpnn_features["mpnn"].to(device)
# RDKit 特征(移到正确设备)
all_features["morgan"] = rdkit_features["morgan"].to(device)
all_features["maccs"] = rdkit_features["maccs"].to(device)
all_features["desc"] = rdkit_features["desc"].to(device)
# Tabular 特征(已在正确设备上)
all_features["comp"] = tabular["comp"]
all_features["phys"] = tabular["phys"]
all_features["help"] = tabular["help"]
all_features["exp"] = tabular["exp"]
# 3. Token Projector: 统一维度
projected = self.token_projector(all_features) # Dict[str, [B, d_model]]
# 4. Stack tokens: [B, n_tokens, d_model]
if self.use_mpnn:
token_order = ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
else:
token_order = ["morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
stacked = torch.stack([projected[k] for k in token_order], dim=1)
projected = self.token_projector(all_features)
stacked = torch.stack([projected[k] for k in self.token_order], dim=1)
return stacked
def _attended_with_moe(self, stacked: torch.Tensor) -> torch.Tensor:
"""
Cross-attention + 可选 MoE 旁路 fusion 输入序列
def _backbone_from_stacked(
self, stacked: torch.Tensor, smiles: Optional[List[str]] = None
) -> torch.Tensor:
chem = stacked[:, : self.split_idx, :]
tab = stacked[:, self.split_idx :, :]
- 不启用 MoE 返回 [B, n_tokens, d]与原行为一致
- 启用 MoE 在最后追加一个 F_moe token返回 [B, n_tokens + 1, d]
chem = self.set_transformer(chem) # chem'
副作用:
把本次 forward MoE 副产物aux loss / gates / probs写到
self._last_moe_extrastrainer 端通过 get_last_moe_extras() 读取
"""
attended = self.cross_attention(stacked)
f_moe = None
if self.moe is not None:
split = self.cross_attention.split_idx
chem_prime = attended[:, :split, :]
tab_prime = attended[:, split:, :]
F_moe, extras = self.moe(chem_prime, tab_prime)
f_moe, extras = self.moe(chem, tab)
self._last_moe_extras = extras
attended = torch.cat([attended, F_moe.unsqueeze(1)], dim=1)
else:
self._last_moe_extras = None
return attended
f_llm = None
if self.llm_prompt is not None and smiles is not None:
f_llm = self.llm_prompt(smiles)
return self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm)
def forward_from_projected(
self,
@ -225,15 +245,13 @@ class LNPModel(nn.Module):
从已投影的 stacked tokens 开始 forward用于 Captum 归因
Args:
stacked: [B, n_tokens, d_model] TokenProjector 输出后 stack 的张量
task: 指定单任务名 ("size", "pdi", "ee", "delivery", "biodist", "toxic")
若为 None返回 delivery head 的标量输出
stacked: [B, n_tokens, d_model]
task: 单任务名None 时返回 delivery head 输出
Returns:
[B, 1] [B, num_classes] 对应任务的预测输出
对应任务的预测输出
"""
attended = self._attended_with_moe(stacked)
fused = self.fusion(attended)
fused = self._backbone_from_stacked(stacked)
if task is None:
task = "delivery"
@ -259,19 +277,10 @@ class LNPModel(nn.Module):
用原始特征替换 base_projected 中指定 token 的投影然后 forward
用于对单个 token 内部特征做 Captum 归因 desc 210
Args:
raw_feature: [B, input_dim] 某个 token 的原始特征
feature_key: token 名称 "desc"
base_projected: [B, n_tokens, d_model] 其他 token 已投影好的张量
task: 任务名
Returns:
对应任务的预测输出
"""
projected = self.token_projector.projectors[feature_key](raw_feature)
gate = torch.sigmoid(self.token_projector.weights[feature_key])
projected = projected * gate # [B, d_model]
projected = projected * gate
token_order = list(self.token_projector.keys)
token_idx = token_order.index(feature_key)
@ -286,38 +295,16 @@ class LNPModel(nn.Module):
smiles: List[str],
tabular: Dict[str, torch.Tensor],
) -> torch.Tensor:
"""
Backbone forward编码 -> 投影 -> 注意力 -> (可选 MoE) -> 融合不经过任务头
用于 pretrain 阶段或需要提取特征的场景
Args:
smiles: SMILES 字符串列表长度为 B
tabular: Dict[str, Tensor]
Returns:
fused: [B, fusion_dim] 融合后的特征向量
"""
"""Backbone forward编码 -> 投影 -> set transformer -> (MoE/LLM) -> 融合。"""
stacked = self._encode_and_project(smiles, tabular)
attended = self._attended_with_moe(stacked)
fused = self.fusion(attended)
return fused
return self._backbone_from_stacked(stacked, smiles=smiles)
def forward_delivery(
self,
smiles: List[str],
tabular: Dict[str, torch.Tensor],
) -> torch.Tensor:
"""
仅预测 delivery用于 pretrain
Args:
smiles: SMILES 字符串列表长度为 B
tabular: Dict[str, Tensor]
Returns:
delivery: [B, 1] 预测的 delivery
"""
"""仅预测 delivery用于 pretrain。返回 [B, 1]。"""
fused = self.forward_backbone(smiles, tabular)
return self.head.delivery_head(fused)
@ -329,52 +316,33 @@ class LNPModel(nn.Module):
"""
完整的多任务 forward
Args:
smiles: SMILES 字符串列表长度为 B
tabular: Dict[str, Tensor]包含:
- "comp": [B, 5] 配方比例
- "phys": [B, 12] 物理参数
- "help": [B, 4] Helper lipid
- "exp": [B, 32] 实验条件
Returns:
Dict[str, Tensor]:
- "size": [B, 1]
- "pdi": [B, 4]
- "ee": [B, 3]
- "delivery": [B, 1]
- "biodist": [B, 7]
- "toxic": [B, 2]
Dict[str, Tensor]: size [B,1], pdi [B,4], ee [B,3],
delivery [B,1], biodist [B,7], toxic [B,2]
"""
fused = self.forward_backbone(smiles, tabular)
outputs = self.head(fused)
return outputs
return self.head(fused)
def clear_cache(self) -> None:
"""清空所有 encoder 的缓存"""
self.rdkit_encoder.clear_cache()
if self.mpnn_encoder is not None:
self.mpnn_encoder.clear_cache()
if self.llm_prompt is not None and hasattr(self.llm_prompt, "clear_cache"):
self.llm_prompt.clear_cache()
def get_last_moe_extras(self) -> Optional[Dict[str, torch.Tensor]]:
"""返回最近一次 forward 中 MoE 模块的副产物aux loss、gates、probs
若未启用 MoE 或还未调用过 forward返回 None
"""
"""返回最近一次 forward 中 MoE 模块的副产物aux loss、gates、probs"""
return self._last_moe_extras
def get_backbone_state_dict(self) -> Dict[str, torch.Tensor]:
"""
获取 backbone 部分的 state_dict不含任务头
包含: token_projector, cross_attention, fusion以及启用时moe
获取 backbone 部分的 state_dict不含任务头且排除冻结的 MolT5 encoder
"""
backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.")
backbone_keys = [
name for name in self.state_dict().keys()
if name.startswith(backbone_prefixes)
]
return {k: v for k, v in self.state_dict().items() if k in backbone_keys}
return {
k: v for k, v in self.state_dict().items()
if k.startswith(BACKBONE_PREFIXES) and not k.startswith(LLM_FROZEN_PREFIX)
}
def get_delivery_head_state_dict(self) -> Dict[str, torch.Tensor]:
"""获取 delivery head 的 state_dict"""
@ -392,16 +360,11 @@ class LNPModel(nn.Module):
"""
从预训练 checkpoint 加载 backbone 可选delivery head 权重
Args:
pretrain_state_dict: 预训练模型的 state_dict
load_delivery_head: 是否加载 delivery head 权重
strict: 是否严格匹配默认 False允许缺失/多余的键
冻结的 MolT5 encoder 权重不在加载范围内
"""
backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.")
keys_to_load = []
for name in pretrain_state_dict.keys():
if name.startswith(backbone_prefixes):
if name.startswith(BACKBONE_PREFIXES) and not name.startswith(LLM_FROZEN_PREFIX):
keys_to_load.append(name)
elif load_delivery_head and name.startswith("head.delivery_head."):
keys_to_load.append(name)
@ -410,45 +373,49 @@ class LNPModel(nn.Module):
k: v for k, v in pretrain_state_dict.items() if k in keys_to_load
}
missing, unexpected = [], []
unexpected = []
model_state = self.state_dict()
for k, v in filtered_state_dict.items():
if k in model_state:
if model_state[k].shape == v.shape:
if k in model_state and model_state[k].shape == v.shape:
model_state[k] = v
else:
unexpected.append(
f"{k} (shape mismatch: {model_state[k].shape} vs {v.shape})"
)
else:
unexpected.append(k)
self.load_state_dict(model_state, strict=False)
if strict and (missing or unexpected):
raise RuntimeError(f"Missing keys: {missing}, Unexpected keys: {unexpected}")
if strict and unexpected:
raise RuntimeError(f"Unexpected keys: {unexpected}")
class LNPModelWithoutMPNN(LNPModel):
"""不使用 MPNN 的简化版本"""
"""不使用 MPNN 的简化版本(化学 token 为 3 个)"""
def __init__(
self,
d_model: int = 256,
num_heads: int = 8,
n_attn_layers: int = 4,
set_transformer_block: str = "sab",
fusion_strategy: PoolingStrategy = "attention",
head_hidden_dim: int = 128,
dropout: float = 0.1,
input_dims: Optional[Dict[str, int]] = None,
# ============ MoE 相关(新增) ============
# ============ MoE 相关 ============
use_moe: bool = False,
moe_n_experts: int = 4,
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
# ============ LLM 相关 ============
use_llm: bool = False,
llm_model_path: str = DEFAULT_MOLT5_PATH,
llm_freeze: bool = True,
llm_use_lora: bool = False,
llm_lora_r: int = 8,
llm_lora_alpha: int = 16,
llm_lora_dropout: float = 0.05,
) -> None:
# 移除 mpnn 维度
dims = input_dims or DEFAULT_INPUT_DIMS.copy()
dims.pop("mpnn", None)
@ -456,6 +423,7 @@ class LNPModelWithoutMPNN(LNPModel):
d_model=d_model,
num_heads=num_heads,
n_attn_layers=n_attn_layers,
set_transformer_block=set_transformer_block,
fusion_strategy=fusion_strategy,
head_hidden_dim=head_hidden_dim,
dropout=dropout,
@ -467,5 +435,11 @@ class LNPModelWithoutMPNN(LNPModel):
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise,
use_llm=use_llm,
llm_model_path=llm_model_path,
llm_freeze=llm_freeze,
llm_use_lora=llm_use_lora,
llm_lora_r=llm_lora_r,
llm_lora_alpha=llm_lora_alpha,
llm_lora_dropout=llm_lora_dropout,
)

View File

@ -44,6 +44,8 @@ from lnp_ml.dataset import (
from tqdm import tqdm
from lnp_ml.modeling.models import LNPModel, LNPModelWithoutMPNN
from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH
from lnp_ml.utils.seed import set_global_seed
from lnp_ml.modeling.encoders.rdkit_encoder import CachedRDKitEncoder
from lnp_ml.modeling.trainer_balanced import (
ClassWeights,
@ -178,19 +180,25 @@ def create_model(
dropout: float = 0.1,
use_mpnn: bool = False,
mpnn_device: str = "cpu",
set_transformer_block: str = "sab",
# MoE
use_moe: bool = False,
moe_n_experts: int = 4,
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
# LLM打包成 dict 以减少多进程参数传递)
llm_kwargs: Optional[Dict] = None,
) -> Union[LNPModel, LNPModelWithoutMPNN]:
"""创建模型"""
moe_kwargs = dict(
"""创建模型。llm_kwargs 含 use_llm / llm_model_path / llm_freeze / llm_use_lora / llm_lora_*。"""
extra_kwargs = dict(
set_transformer_block=set_transformer_block,
use_moe=use_moe,
moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise,
**(llm_kwargs or {}),
)
if use_mpnn:
@ -204,7 +212,7 @@ def create_model(
dropout=dropout,
mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device,
**moe_kwargs,
**extra_kwargs,
)
else:
return LNPModelWithoutMPNN(
@ -214,7 +222,7 @@ def create_model(
fusion_strategy=fusion_strategy,
head_hidden_dim=head_hidden_dim,
dropout=dropout,
**moe_kwargs,
**extra_kwargs,
)
@ -381,6 +389,8 @@ def run_inner_optuna(
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
set_transformer_block: str = "sab",
llm_kwargs: Optional[Dict] = None,
) -> Tuple[Dict, int, optuna.Study]:
"""
在内层数据上运行 Optuna 超参搜索
@ -436,6 +446,16 @@ def run_inner_optuna(
weight_decay = trial.suggest_float("weight_decay", 1e-5, 1e-1, log=True)
backbone_lr_ratio = trial.suggest_float("backbone_lr_ratio", 0.01, 1.0, log=True)
# MoE / LLM
use_llm = bool(llm_kwargs.get("use_llm", False))
moe_ne_t = trial.suggest_categorical("moe_n_experts", [2, 4]) if use_moe else moe_n_experts
moe_tk_t = trial.suggest_int("moe_top_k", 1, 2) if use_moe else moe_top_k
moe_hm_t = trial.suggest_categorical("moe_expert_hidden_mult", [1, 2]) if use_moe else moe_expert_hidden_mult
llm_kwargs_t = dict(llm_kwargs)
if use_llm:
llm_kwargs_t["llm_use_lora"] = True
llm_kwargs_t["llm_lora_r"] = trial.suggest_categorical("llm_lora_r", [8, 16, 32])
# 内层 3-fold CV
inner_cv = StratifiedKFold(
n_splits=n_inner_folds, shuffle=True, random_state=seed
@ -476,10 +496,12 @@ def run_inner_optuna(
use_mpnn=use_mpnn,
mpnn_device=device.type,
use_moe=use_moe,
moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_n_experts=moe_ne_t,
moe_top_k=moe_tk_t,
moe_expert_hidden_mult=moe_hm_t,
moe_jitter_noise=moe_jitter_noise,
set_transformer_block=set_transformer_block,
llm_kwargs=llm_kwargs_t,
)
if rdkit_cache is not None:
model.rdkit_encoder._cache = rdkit_cache
@ -536,6 +558,7 @@ def run_inner_optuna(
"n_attn_layers": fixed_n_attn_layers,
"fusion_strategy": fixed_fusion_strategy,
"head_hidden_dim": fixed_head_hidden_dim,
"set_transformer_block": set_transformer_block,
})
epoch_mean = study.best_trial.user_attrs.get("epoch_mean", epochs_per_trial)
@ -571,6 +594,10 @@ def _run_single_outer_fold(
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
set_transformer_block: str = "sab",
llm_kwargs: Optional[Dict] = None,
precomputed_best_params: Optional[Dict] = None,
precomputed_epoch_mean: Optional[int] = None,
) -> Dict:
"""
执行单个外层 fold 的完整流程内层调参 + 外层训练 + 评估
@ -578,6 +605,7 @@ def _run_single_outer_fold(
所有参数均为可序列化类型以支持 spawn 多进程
"""
device = torch.device(device_str)
set_global_seed(seed + outer_fold)
fold_dir = Path(fold_dir)
fold_dir.mkdir(parents=True, exist_ok=True)
@ -598,10 +626,14 @@ def _run_single_outer_fold(
"outer_test_idx": outer_test_idx.tolist(),
}, f)
# 内层 Optuna 调参
# 内层 Optuna 调参(方式 B已有超参则跳过直接复用
if precomputed_best_params is not None and precomputed_epoch_mean is not None:
logger.info("Reusing precomputed best_params (skip inner Optuna).")
best_params = dict(precomputed_best_params)
epoch_mean = int(precomputed_epoch_mean)
else:
logger.info(f"\nRunning inner Optuna with {n_trials} trials...")
study_path = fold_dir / "optuna_study.sqlite3"
best_params, epoch_mean, study = run_inner_optuna(
full_dataset=full_dataset,
inner_train_indices=outer_train_idx,
@ -624,24 +656,36 @@ def _run_single_outer_fold(
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise,
set_transformer_block=set_transformer_block,
llm_kwargs=llm_kwargs,
)
# 保存最佳参数
with open(fold_dir / "best_params.json", "w") as f:
json.dump(best_params, f, indent=2)
with open(fold_dir / "epoch_mean.json", "w") as f:
json.dump({"epoch_mean": epoch_mean}, f)
# 外层训练(使用最优超参,固定 epoch 数,不 early-stop
logger.info(f"\nTraining outer fold with best params, epochs={epoch_mean}...")
train_subset = Subset(full_dataset, outer_train_idx.tolist())
# 从 outer_train 再切 10% 作为"监控用 val"(只画曲线、看 gap不参与选模型/早停)
_rng = np.random.RandomState(seed + outer_fold)
_perm = _rng.permutation(len(outer_train_idx))
_n_val = max(1, int(0.1 * len(outer_train_idx)))
fit_idx = outer_train_idx[_perm[_n_val:]]
monitor_idx = outer_train_idx[_perm[:_n_val]]
train_subset = Subset(full_dataset, fit_idx.tolist())
monitor_subset = Subset(full_dataset, monitor_idx.tolist())
test_subset = Subset(full_dataset, outer_test_idx.tolist())
train_loader = DataLoader(
train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn
)
monitor_loader = DataLoader(
monitor_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn
)
test_loader = DataLoader(
test_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn
)
@ -658,10 +702,13 @@ def _run_single_outer_fold(
use_mpnn=use_mpnn,
mpnn_device=device.type,
use_moe=use_moe,
moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_n_experts=best_params.get("moe_n_experts", moe_n_experts),
moe_top_k=best_params.get("moe_top_k", moe_top_k),
moe_expert_hidden_mult=best_params.get("moe_expert_hidden_mult", moe_expert_hidden_mult),
moe_jitter_noise=moe_jitter_noise,
set_transformer_block=best_params.get("set_transformer_block", "sab"),
llm_kwargs={**llm_kwargs, **({"llm_use_lora": True, "llm_lora_r": best_params["llm_lora_r"]}
if (llm_kwargs.get("use_llm", False) and "llm_lora_r" in best_params) else {})},
)
model.rdkit_encoder._cache = rdkit_cache
@ -676,7 +723,7 @@ def _run_single_outer_fold(
train_result = train_fixed_epochs(
model=model,
train_loader=train_loader,
val_loader=None,
val_loader=monitor_loader,
device=device,
lr=best_params["lr"],
weight_decay=best_params["weight_decay"],
@ -684,6 +731,7 @@ def _run_single_outer_fold(
class_weights=class_weights,
use_cosine_annealing=True,
backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0),
freeze_backbone_epochs=3,
)
model.load_state_dict(train_result["final_state"])
@ -695,6 +743,7 @@ def _run_single_outer_fold(
"n_attn_layers": best_params["n_attn_layers"],
"fusion_strategy": best_params["fusion_strategy"],
"head_hidden_dim": best_params["head_hidden_dim"],
"set_transformer_block": best_params.get("set_transformer_block", "sab"),
"dropout": best_params["dropout"],
"use_mpnn": use_mpnn,
"use_moe": use_moe,
@ -702,6 +751,7 @@ def _run_single_outer_fold(
"moe_top_k": moe_top_k,
"moe_expert_hidden_mult": moe_expert_hidden_mult,
"moe_jitter_noise": moe_jitter_noise,
**(llm_kwargs or {}),
}
torch.save({
@ -759,12 +809,24 @@ def main(
load_delivery_head: bool = False,
# MPNN
use_mpnn: bool = False,
# MoE新增
n_repeats: int = 1,
repeat_seed_step: int = 1000,
# MoE消融开关
use_moe: bool = False,
moe_n_experts: int = 4,
moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0,
# Set Transformer
set_transformer_block: str = "sab",
# LLM消融开关
use_llm: bool = False,
llm_model_path: str = DEFAULT_MOLT5_PATH,
llm_freeze: bool = True,
llm_use_lora: bool = False,
llm_lora_r: int = 8,
llm_lora_alpha: int = 16,
llm_lora_dropout: float = 0.05,
# 并行
parallel: bool = False,
# 设备
@ -785,6 +847,17 @@ def main(
logger.info(f"Using device: {device}")
device = torch.device(device)
set_global_seed(seed)
llm_kwargs = dict(
use_llm=use_llm,
llm_model_path=llm_model_path,
llm_freeze=llm_freeze,
llm_use_lora=llm_use_lora,
llm_lora_r=llm_lora_r,
llm_lora_alpha=llm_lora_alpha,
llm_lora_dropout=llm_lora_dropout,
)
# 加载预训练权重(如果指定)
pretrain_state_dict = None
@ -858,6 +931,8 @@ def main(
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise,
set_transformer_block=set_transformer_block,
llm_kwargs=llm_kwargs,
))
if parallel:
@ -873,9 +948,22 @@ def main(
else:
logger.info(f"Running {n_outer_folds} outer folds SEQUENTIALLY")
outer_results = []
per_fold_params = {}
for args in fold_args:
result = _run_single_outer_fold(**args)
outer_results.append(result)
per_fold_params[result["fold"]] = (result["best_params"], result["epoch_mean"])
for rep in range(1, n_repeats):
logger.info(f"\n===== Repeat {rep}/{n_repeats - 1} (reuse hyperparams) =====")
for args in fold_args:
bp, em = per_fold_params[args["outer_fold"]]
rep_args = dict(args)
rep_args["seed"] = seed + rep * repeat_seed_step
rep_args["fold_dir"] = run_dir / f"repeat_{rep}" / f"outer_fold_{args['outer_fold']}"
rep_args["precomputed_best_params"] = bp
rep_args["precomputed_epoch_mean"] = em
outer_results.append(_run_single_outer_fold(**rep_args))
# 汇总结果
logger.info("\n" + "=" * 60)

View File

@ -55,6 +55,13 @@ def load_model(
fusion_strategy=config["fusion_strategy"],
head_hidden_dim=config["head_hidden_dim"],
dropout=config["dropout"],
use_llm=config.get("use_llm", False),
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
llm_freeze=config.get("llm_freeze", True),
llm_use_lora=config.get("llm_use_lora", False),
llm_lora_r=config.get("llm_lora_r", 8),
llm_lora_alpha=config.get("llm_lora_alpha", 16),
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device,
)
@ -66,9 +73,16 @@ def load_model(
fusion_strategy=config["fusion_strategy"],
head_hidden_dim=config["head_hidden_dim"],
dropout=config["dropout"],
use_llm=config.get("use_llm", False),
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
llm_freeze=config.get("llm_freeze", True),
llm_use_lora=config.get("llm_use_lora", False),
llm_lora_r=config.get("llm_lora_r", 8),
llm_lora_alpha=config.get("llm_lora_alpha", 16),
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
)
model.load_state_dict(checkpoint["model_state_dict"])
model.load_state_dict(checkpoint["model_state_dict"], strict=False)
model.to(device)
model.eval()