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 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 tabexpert 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_extrastrainer 端通过 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

View File

@ -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()