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