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 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 看 tab,expert 吃 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_extras,trainer 端通过 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,
|
||||
)
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user