mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +08:00
Add Mixture-of-Experts (MoE) layer support across modeling pipeline
This commit is contained in:
parent
804bda5576
commit
cd4534f05a
16
Makefile
16
Makefile
@ -18,6 +18,14 @@ USE_SWA_FLAG = $(if $(USE_SWA),--use-swa,)
|
|||||||
PARALLEL_FLAG = $(if $(PARALLEL),--parallel,)
|
PARALLEL_FLAG = $(if $(PARALLEL),--parallel,)
|
||||||
INIT_PRETRAIN_FLAG = $(if $(NO_PRETRAIN),,--init-from-pretrain $(or $(INIT_PRETRAIN),models/pretrain_delivery.pt))
|
INIT_PRETRAIN_FLAG = $(if $(NO_PRETRAIN),,--init-from-pretrain $(or $(INIT_PRETRAIN),models/pretrain_delivery.pt))
|
||||||
|
|
||||||
|
# --- MoE 相关 flag ---
|
||||||
|
USE_MOE_FLAG = $(if $(USE_MOE),--use-moe,)
|
||||||
|
MOE_EXPERTS_FLAG = $(if $(MOE_EXPERTS),--moe-n-experts $(MOE_EXPERTS),)
|
||||||
|
MOE_TOPK_FLAG = $(if $(MOE_TOPK),--moe-top-k $(MOE_TOPK),)
|
||||||
|
MOE_HIDDEN_FLAG = $(if $(MOE_HIDDEN_MULT),--moe-expert-hidden-mult $(MOE_HIDDEN_MULT),)
|
||||||
|
MOE_JITTER_FLAG = $(if $(MOE_JITTER),--moe-jitter-noise $(MOE_JITTER),)
|
||||||
|
MOE_FLAGS = $(USE_MOE_FLAG) $(MOE_EXPERTS_FLAG) $(MOE_TOPK_FLAG) $(MOE_HIDDEN_FLAG) $(MOE_JITTER_FLAG)
|
||||||
|
|
||||||
#################################################################################
|
#################################################################################
|
||||||
# ENVIRONMENT & CODE QUALITY #
|
# ENVIRONMENT & CODE QUALITY #
|
||||||
#################################################################################
|
#################################################################################
|
||||||
@ -117,11 +125,13 @@ pretrain: requirements
|
|||||||
train: requirements
|
train: requirements
|
||||||
$(PYTHON_INTERPRETER) -m lnp_ml.modeling.nested_cv_optuna \
|
$(PYTHON_INTERPRETER) -m lnp_ml.modeling.nested_cv_optuna \
|
||||||
$(DEVICE_FLAG) $(MPNN_FLAG) $(SEED_FLAG) $(INIT_PRETRAIN_FLAG) \
|
$(DEVICE_FLAG) $(MPNN_FLAG) $(SEED_FLAG) $(INIT_PRETRAIN_FLAG) \
|
||||||
$(N_TRIALS_FLAG) $(EPOCHS_PER_TRIAL_FLAG) $(MIN_STRATUM_FLAG) $(OUTPUT_DIR_FLAG) $(PARALLEL_FLAG)
|
$(N_TRIALS_FLAG) $(EPOCHS_PER_TRIAL_FLAG) $(MIN_STRATUM_FLAG) $(OUTPUT_DIR_FLAG) $(PARALLEL_FLAG) \
|
||||||
|
$(MOE_FLAGS)
|
||||||
$(PYTHON_INTERPRETER) -m lnp_ml.modeling.final_train_optuna_cv \
|
$(PYTHON_INTERPRETER) -m lnp_ml.modeling.final_train_optuna_cv \
|
||||||
$(DEVICE_FLAG) $(MPNN_FLAG) $(SEED_FLAG) $(INIT_PRETRAIN_FLAG) \
|
$(DEVICE_FLAG) $(MPNN_FLAG) $(SEED_FLAG) $(INIT_PRETRAIN_FLAG) \
|
||||||
$(N_TRIALS_FLAG) $(EPOCHS_PER_TRIAL_FLAG) $(MIN_STRATUM_FLAG) $(OUTPUT_DIR_FLAG) $(USE_SWA_FLAG)
|
$(N_TRIALS_FLAG) $(EPOCHS_PER_TRIAL_FLAG) $(MIN_STRATUM_FLAG) $(OUTPUT_DIR_FLAG) $(USE_SWA_FLAG) \
|
||||||
|
$(MOE_FLAGS)
|
||||||
|
|
||||||
#################################################################################
|
#################################################################################
|
||||||
# INTERPRETABILITY (biodistribution feature importance) #
|
# INTERPRETABILITY (biodistribution feature importance) #
|
||||||
#################################################################################
|
#################################################################################
|
||||||
|
|||||||
@ -41,8 +41,12 @@ BIODIST_ORGANS = ["lymph_nodes", "heart", "liver", "spleen", "lung", "kidney", "
|
|||||||
|
|
||||||
def get_token_names(model: Union[LNPModel, LNPModelWithoutMPNN]) -> List[str]:
|
def get_token_names(model: Union[LNPModel, LNPModelWithoutMPNN]) -> List[str]:
|
||||||
if model.use_mpnn:
|
if model.use_mpnn:
|
||||||
return ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
|
names = ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
|
||||||
return ["morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
|
else:
|
||||||
|
names = ["morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
|
||||||
|
if getattr(model, "moe", None) is not None:
|
||||||
|
names.append("moe")
|
||||||
|
return names
|
||||||
|
|
||||||
|
|
||||||
# ──────────────────────────────────────────────────────────────────────
|
# ──────────────────────────────────────────────────────────────────────
|
||||||
@ -192,7 +196,8 @@ def fusion_attention_importance(
|
|||||||
all_weights: List[torch.Tensor] = []
|
all_weights: List[torch.Tensor] = []
|
||||||
for start in range(0, all_tokens.size(0), batch_size):
|
for start in range(0, all_tokens.size(0), batch_size):
|
||||||
end = min(start + batch_size, all_tokens.size(0))
|
end = min(start + batch_size, all_tokens.size(0))
|
||||||
attended = model.cross_attention(all_tokens[start:end].to(device))
|
# 使用与 forward_from_projected 相同的桥梁,自动处理 MoE 的额外 token
|
||||||
|
attended = model._attended_with_moe(all_tokens[start:end].to(device))
|
||||||
_, weights = model.fusion(attended, return_attn_weights=True)
|
_, weights = model.fusion(attended, return_attn_weights=True)
|
||||||
all_weights.append(weights.cpu())
|
all_weights.append(weights.cpu())
|
||||||
|
|
||||||
@ -295,7 +300,7 @@ def plot_token_importance(
|
|||||||
vals_sorted = normed[order]
|
vals_sorted = normed[order]
|
||||||
|
|
||||||
n_tokens = len(token_names)
|
n_tokens = len(token_names)
|
||||||
split_idx = 4 if n_tokens == 8 else 3
|
split_idx = 4 if "mpnn" in token_names else 3
|
||||||
channel_a_set = set(token_names[:split_idx])
|
channel_a_set = set(token_names[:split_idx])
|
||||||
colors = [color_a if n in channel_a_set else color_b for n in names_sorted]
|
colors = [color_a if n in channel_a_set else color_b for n in names_sorted]
|
||||||
|
|
||||||
|
|||||||
@ -176,8 +176,22 @@ def create_model(
|
|||||||
dropout: float = 0.1,
|
dropout: float = 0.1,
|
||||||
use_mpnn: bool = False,
|
use_mpnn: bool = False,
|
||||||
mpnn_device: str = "cpu",
|
mpnn_device: str = "cpu",
|
||||||
|
# ============ 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,
|
||||||
) -> Union[LNPModel, LNPModelWithoutMPNN]:
|
) -> Union[LNPModel, LNPModelWithoutMPNN]:
|
||||||
"""创建模型"""
|
"""创建模型"""
|
||||||
|
moe_kwargs = dict(
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
if use_mpnn:
|
if use_mpnn:
|
||||||
ensemble_paths = find_mpnn_ensemble_paths()
|
ensemble_paths = find_mpnn_ensemble_paths()
|
||||||
return LNPModel(
|
return LNPModel(
|
||||||
@ -189,6 +203,7 @@ def create_model(
|
|||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
mpnn_ensemble_paths=ensemble_paths,
|
mpnn_ensemble_paths=ensemble_paths,
|
||||||
mpnn_device=mpnn_device,
|
mpnn_device=mpnn_device,
|
||||||
|
**moe_kwargs,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return LNPModelWithoutMPNN(
|
return LNPModelWithoutMPNN(
|
||||||
@ -198,6 +213,7 @@ def create_model(
|
|||||||
fusion_strategy=fusion_strategy,
|
fusion_strategy=fusion_strategy,
|
||||||
head_hidden_dim=head_hidden_dim,
|
head_hidden_dim=head_hidden_dim,
|
||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
|
**moe_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@ -232,7 +248,6 @@ def load_pretrain_weights_to_model(
|
|||||||
|
|
||||||
|
|
||||||
# ============ 3-fold Optuna 调参 ============
|
# ============ 3-fold Optuna 调参 ============
|
||||||
|
|
||||||
def run_optuna_cv(
|
def run_optuna_cv(
|
||||||
full_dataset: LNPDataset,
|
full_dataset: LNPDataset,
|
||||||
strata: np.ndarray,
|
strata: np.ndarray,
|
||||||
@ -249,6 +264,11 @@ def run_optuna_cv(
|
|||||||
pretrain_config: Optional[Dict] = None,
|
pretrain_config: Optional[Dict] = None,
|
||||||
load_delivery_head: bool = True,
|
load_delivery_head: bool = True,
|
||||||
rdkit_cache: Optional[Dict] = None,
|
rdkit_cache: Optional[Dict] = None,
|
||||||
|
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,
|
||||||
) -> Tuple[Dict, int, optuna.Study]:
|
) -> Tuple[Dict, int, optuna.Study]:
|
||||||
"""
|
"""
|
||||||
使用全量数据做 3-fold CV Optuna 超参搜索。
|
使用全量数据做 3-fold CV Optuna 超参搜索。
|
||||||
@ -335,6 +355,11 @@ def run_optuna_cv(
|
|||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
use_mpnn=use_mpnn,
|
use_mpnn=use_mpnn,
|
||||||
mpnn_device=device.type,
|
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_jitter_noise=moe_jitter_noise,
|
||||||
)
|
)
|
||||||
if rdkit_cache is not None:
|
if rdkit_cache is not None:
|
||||||
model.rdkit_encoder._cache = rdkit_cache
|
model.rdkit_encoder._cache = rdkit_cache
|
||||||
@ -404,7 +429,6 @@ def run_optuna_cv(
|
|||||||
|
|
||||||
|
|
||||||
# ============ 主流程 ============
|
# ============ 主流程 ============
|
||||||
|
|
||||||
@app.command()
|
@app.command()
|
||||||
def main(
|
def main(
|
||||||
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
|
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
|
||||||
@ -427,6 +451,12 @@ def main(
|
|||||||
load_delivery_head: bool = False,
|
load_delivery_head: bool = False,
|
||||||
# MPNN
|
# MPNN
|
||||||
use_mpnn: bool = False,
|
use_mpnn: bool = False,
|
||||||
|
# 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,
|
||||||
# 设备
|
# 设备
|
||||||
device: str = "cuda" if torch.cuda.is_available() else "cpu",
|
device: str = "cuda" if torch.cuda.is_available() else "cpu",
|
||||||
):
|
):
|
||||||
@ -507,6 +537,11 @@ def main(
|
|||||||
pretrain_config=pretrain_config,
|
pretrain_config=pretrain_config,
|
||||||
load_delivery_head=load_delivery_head,
|
load_delivery_head=load_delivery_head,
|
||||||
rdkit_cache=rdkit_cache,
|
rdkit_cache=rdkit_cache,
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 保存最佳参数
|
# 保存最佳参数
|
||||||
@ -564,6 +599,11 @@ def main(
|
|||||||
dropout=best_params["dropout"],
|
dropout=best_params["dropout"],
|
||||||
use_mpnn=use_mpnn,
|
use_mpnn=use_mpnn,
|
||||||
mpnn_device=device.type,
|
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_jitter_noise=moe_jitter_noise,
|
||||||
)
|
)
|
||||||
model.rdkit_encoder._cache = rdkit_cache
|
model.rdkit_encoder._cache = rdkit_cache
|
||||||
|
|
||||||
@ -611,8 +651,13 @@ def main(
|
|||||||
"head_hidden_dim": best_params["head_hidden_dim"],
|
"head_hidden_dim": best_params["head_hidden_dim"],
|
||||||
"dropout": best_params["dropout"],
|
"dropout": best_params["dropout"],
|
||||||
"use_mpnn": use_mpnn,
|
"use_mpnn": use_mpnn,
|
||||||
|
"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,
|
||||||
}
|
}
|
||||||
|
|
||||||
torch.save({
|
torch.save({
|
||||||
"model_state_dict": train_result["final_state"],
|
"model_state_dict": train_result["final_state"],
|
||||||
"config": config,
|
"config": config,
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
from lnp_ml.modeling.layers.token_projector import TokenProjector
|
from lnp_ml.modeling.layers.token_projector import TokenProjector
|
||||||
from lnp_ml.modeling.layers.bidirectional_cross_attention import CrossModalAttention
|
from lnp_ml.modeling.layers.bidirectional_cross_attention import CrossModalAttention
|
||||||
from lnp_ml.modeling.layers.fusion import FusionLayer
|
from lnp_ml.modeling.layers.fusion import FusionLayer
|
||||||
|
from lnp_ml.modeling.layers.moe import MoEBlock
|
||||||
|
|
||||||
__all__ = ["TokenProjector", "CrossModalAttention", "FusionLayer"]
|
__all__ = ["TokenProjector", "CrossModalAttention", "FusionLayer", "MoEBlock"]
|
||||||
|
|
||||||
229
lnp_ml/modeling/layers/moe.py
Normal file
229
lnp_ml/modeling/layers/moe.py
Normal file
@ -0,0 +1,229 @@
|
|||||||
|
"""
|
||||||
|
MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。
|
||||||
|
|
||||||
|
设计要点(与现有 8-token 架构对齐):
|
||||||
|
- Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。
|
||||||
|
- Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。
|
||||||
|
- 路由粒度:每个样本只算一次 router → gates [B, K]。
|
||||||
|
- Top-k 稀疏激活 + 可选训练态 jitter 噪声。
|
||||||
|
- 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。
|
||||||
|
|
||||||
|
输出:
|
||||||
|
F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。
|
||||||
|
extras: dict 诊断与监控信息(aux loss、gates 等)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Dict, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
class MoEAttentionPool(nn.Module):
|
||||||
|
"""与 FusionLayer 同款的 attention pooling,参数独立。
|
||||||
|
|
||||||
|
将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, d_model: int) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.d_model = d_model
|
||||||
|
self.query = nn.Parameter(torch.randn(1, 1, d_model))
|
||||||
|
self.proj = nn.Linear(d_model, d_model)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
x: [B, T, d_model]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[B, d_model]
|
||||||
|
"""
|
||||||
|
B = x.size(0)
|
||||||
|
q = self.query.expand(B, -1, -1) # [B, 1, d]
|
||||||
|
k = self.proj(x) # [B, T, d]
|
||||||
|
scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5)
|
||||||
|
weights = F.softmax(scores, dim=-1) # [B, 1, T]
|
||||||
|
return torch.bmm(weights, x).squeeze(1) # [B, d]
|
||||||
|
|
||||||
|
|
||||||
|
class MoERouter(nn.Module):
|
||||||
|
"""单层 softmax router + Top-k 稀疏激活。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
d_model: router 输入维度。
|
||||||
|
n_experts: 专家数量 K。
|
||||||
|
top_k: 每个样本激活的专家数。
|
||||||
|
jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
d_model: int,
|
||||||
|
n_experts: int,
|
||||||
|
top_k: int = 2,
|
||||||
|
jitter_noise: float = 0.0,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
if not 1 <= top_k <= n_experts:
|
||||||
|
raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}")
|
||||||
|
self.n_experts = n_experts
|
||||||
|
self.top_k = top_k
|
||||||
|
self.jitter_noise = jitter_noise
|
||||||
|
self.linear = nn.Linear(d_model, n_experts)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self, q: torch.Tensor,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
q: [B, d_model] 路由查询向量。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
gates: [B, n_experts] top-k 后再归一化的稀疏概率。
|
||||||
|
probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。
|
||||||
|
expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。
|
||||||
|
"""
|
||||||
|
logits = self.linear(q) # [B, K]
|
||||||
|
if self.training and self.jitter_noise > 0.0:
|
||||||
|
noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise
|
||||||
|
logits = logits + noise
|
||||||
|
|
||||||
|
probs_full = F.softmax(logits, dim=-1) # [B, K]
|
||||||
|
|
||||||
|
topk_vals, topk_idx = probs_full.topk(self.top_k, dim=-1) # [B, k]
|
||||||
|
expert_mask = torch.zeros_like(probs_full)
|
||||||
|
expert_mask.scatter_(1, topk_idx, 1.0)
|
||||||
|
|
||||||
|
gates = probs_full * expert_mask
|
||||||
|
gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9)
|
||||||
|
return gates, probs_full, expert_mask
|
||||||
|
|
||||||
|
|
||||||
|
class MoEExpert(nn.Module):
|
||||||
|
"""单个专家 MLP:in_dim → hidden_dim → out_dim。"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.net = nn.Sequential(
|
||||||
|
nn.Linear(in_dim, hidden_dim),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(hidden_dim, out_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return self.net(x)
|
||||||
|
|
||||||
|
|
||||||
|
class MoEBlock(nn.Module):
|
||||||
|
"""
|
||||||
|
Sample-level 跨模态 MoE。
|
||||||
|
|
||||||
|
流程:
|
||||||
|
1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d]
|
||||||
|
2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask)
|
||||||
|
3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d]
|
||||||
|
4. K 个 Expert MLP 并行处理 -> stack [B, K, d]
|
||||||
|
5. 加权求和 -> F_moe [B, d]
|
||||||
|
6. 计算 load-balancing aux loss
|
||||||
|
|
||||||
|
Args:
|
||||||
|
d_model: token 维度。
|
||||||
|
n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。
|
||||||
|
n_experts: 专家数量 K。
|
||||||
|
top_k: 每个样本激活的专家数。
|
||||||
|
expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。
|
||||||
|
dropout: expert 内部 dropout。
|
||||||
|
jitter_noise: router 训练态噪声幅度,0.0 表示关闭。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
d_model: int,
|
||||||
|
n_chem_tokens: int = 4,
|
||||||
|
n_experts: int = 4,
|
||||||
|
top_k: int = 2,
|
||||||
|
expert_hidden_mult: int = 2,
|
||||||
|
dropout: float = 0.1,
|
||||||
|
jitter_noise: float = 0.0,
|
||||||
|
) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.d_model = d_model
|
||||||
|
self.n_chem_tokens = n_chem_tokens
|
||||||
|
self.n_experts = n_experts
|
||||||
|
self.top_k = top_k
|
||||||
|
|
||||||
|
self.tab_pool = MoEAttentionPool(d_model)
|
||||||
|
self.router = MoERouter(
|
||||||
|
d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise,
|
||||||
|
)
|
||||||
|
|
||||||
|
expert_in = d_model * n_chem_tokens
|
||||||
|
expert_hidden = d_model * expert_hidden_mult
|
||||||
|
self.experts = nn.ModuleList([
|
||||||
|
MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout)
|
||||||
|
for _ in range(n_experts)
|
||||||
|
])
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
chem: torch.Tensor,
|
||||||
|
tab: torch.Tensor,
|
||||||
|
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
chem: [B, T_chem, d_model] 化学侧 token(router 不看)。
|
||||||
|
tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
F_moe: [B, d_model]
|
||||||
|
extras: {
|
||||||
|
"lb_loss": 标量 load-balancing aux loss(带梯度),
|
||||||
|
"gates": [B, K] 稀疏归一后的门控(detached),
|
||||||
|
"probs": [B, K] 原始 softmax 概率(detached),
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
if chem.size(1) != self.n_chem_tokens:
|
||||||
|
raise ValueError(
|
||||||
|
f"chem token 数不匹配:期望 {self.n_chem_tokens}, "
|
||||||
|
f"实际 {chem.size(1)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
q_tab = self.tab_pool(tab) # [B, d]
|
||||||
|
gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K]
|
||||||
|
|
||||||
|
flat = chem.flatten(start_dim=1) # [B, T_chem * d]
|
||||||
|
expert_outs = torch.stack(
|
||||||
|
[expert(flat) for expert in self.experts], dim=1,
|
||||||
|
) # [B, K, d]
|
||||||
|
|
||||||
|
F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d]
|
||||||
|
|
||||||
|
lb_loss = self._load_balancing_loss(probs_full, expert_mask)
|
||||||
|
|
||||||
|
return F_moe, {
|
||||||
|
"lb_loss": lb_loss,
|
||||||
|
"gates": gates.detach(),
|
||||||
|
"probs": probs_full.detach(),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _load_balancing_loss(
|
||||||
|
self,
|
||||||
|
probs_full: torch.Tensor,
|
||||||
|
expert_mask: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Switch Transformer 风格的 load-balancing loss。
|
||||||
|
|
||||||
|
f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值)
|
||||||
|
p_i = 该 batch 中 expert_i 的 softmax 概率均值
|
||||||
|
loss = K * Σ_i f_i * p_i
|
||||||
|
|
||||||
|
理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。
|
||||||
|
"""
|
||||||
|
f = expert_mask.mean(dim=0) # [K]
|
||||||
|
p = probs_full.mean(dim=0) # [K]
|
||||||
|
return self.n_experts * (f * p).sum()
|
||||||
@ -5,7 +5,7 @@ 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
|
from lnp_ml.modeling.layers import TokenProjector, CrossModalAttention, FusionLayer, MoEBlock
|
||||||
from lnp_ml.modeling.heads import MultiTaskHead
|
from lnp_ml.modeling.heads import MultiTaskHead
|
||||||
|
|
||||||
|
|
||||||
@ -62,6 +62,12 @@ class LNPModel(nn.Module):
|
|||||||
mpnn_device: str = "cpu",
|
mpnn_device: str = "cpu",
|
||||||
# 输入维度配置
|
# 输入维度配置
|
||||||
input_dims: Optional[Dict[str, int]] = None,
|
input_dims: Optional[Dict[str, int]] = None,
|
||||||
|
# ============ 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,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@ -107,10 +113,28 @@ class LNPModel(nn.Module):
|
|||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ============ MoE Block (可选) ============
|
||||||
|
self.use_moe = use_moe
|
||||||
|
self._last_moe_extras: Optional[Dict[str, torch.Tensor]] = None
|
||||||
|
if use_moe:
|
||||||
|
self.moe = MoEBlock(
|
||||||
|
d_model=d_model,
|
||||||
|
n_chem_tokens=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 ============
|
# ============ Fusion Layer ============
|
||||||
self.fusion = FusionLayer(
|
self.fusion = FusionLayer(
|
||||||
d_model=d_model,
|
d_model=d_model,
|
||||||
n_tokens=n_tokens,
|
n_tokens=n_fusion_tokens,
|
||||||
strategy=fusion_strategy,
|
strategy=fusion_strategy,
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -169,6 +193,29 @@ class LNPModel(nn.Module):
|
|||||||
stacked = torch.stack([projected[k] for k in token_order], dim=1)
|
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:
|
||||||
|
"""
|
||||||
|
Cross-attention + 可选 MoE 旁路 → fusion 输入序列。
|
||||||
|
|
||||||
|
- 不启用 MoE 时:返回 [B, n_tokens, d],与原行为一致。
|
||||||
|
- 启用 MoE 时:在最后追加一个 F_moe token,返回 [B, n_tokens + 1, d]。
|
||||||
|
|
||||||
|
副作用:
|
||||||
|
把本次 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:
|
||||||
|
split = self.cross_attention.split_idx
|
||||||
|
chem_prime = attended[:, :split, :]
|
||||||
|
tab_prime = attended[:, split:, :]
|
||||||
|
F_moe, extras = self.moe(chem_prime, tab_prime)
|
||||||
|
self._last_moe_extras = extras
|
||||||
|
attended = torch.cat([attended, F_moe.unsqueeze(1)], dim=1)
|
||||||
|
else:
|
||||||
|
self._last_moe_extras = None
|
||||||
|
return attended
|
||||||
|
|
||||||
def forward_from_projected(
|
def forward_from_projected(
|
||||||
self,
|
self,
|
||||||
stacked: torch.Tensor,
|
stacked: torch.Tensor,
|
||||||
@ -178,14 +225,14 @@ 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] TokenProjector 输出后 stack 的张量。
|
||||||
task: 指定单任务名 ("size", "pdi", "ee", "delivery", "biodist", "toxic")。
|
task: 指定单任务名 ("size", "pdi", "ee", "delivery", "biodist", "toxic")。
|
||||||
若为 None,返回 delivery head 的标量输出。
|
若为 None,返回 delivery head 的标量输出。
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
[B, 1] 或 [B, num_classes] 对应任务的预测输出
|
[B, 1] 或 [B, num_classes] 对应任务的预测输出。
|
||||||
"""
|
"""
|
||||||
attended = self.cross_attention(stacked)
|
attended = self._attended_with_moe(stacked)
|
||||||
fused = self.fusion(attended)
|
fused = self.fusion(attended)
|
||||||
|
|
||||||
if task is None:
|
if task is None:
|
||||||
@ -240,26 +287,20 @@ class LNPModel(nn.Module):
|
|||||||
tabular: Dict[str, torch.Tensor],
|
tabular: Dict[str, torch.Tensor],
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
Backbone forward:编码 -> 投影 -> 注意力 -> 融合,不经过任务头。
|
Backbone forward:编码 -> 投影 -> 注意力 -> (可选 MoE) -> 融合,不经过任务头。
|
||||||
|
|
||||||
用于 pretrain 阶段或需要提取特征的场景。
|
用于 pretrain 阶段或需要提取特征的场景。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
smiles: SMILES 字符串列表,长度为 B
|
smiles: SMILES 字符串列表,长度为 B
|
||||||
tabular: Dict[str, Tensor]
|
tabular: Dict[str, Tensor]
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
fused: [B, fusion_dim] 融合后的特征向量
|
fused: [B, fusion_dim] 融合后的特征向量
|
||||||
"""
|
"""
|
||||||
# 编码 + 投影 + stack
|
|
||||||
stacked = self._encode_and_project(smiles, tabular)
|
stacked = self._encode_and_project(smiles, tabular)
|
||||||
|
attended = self._attended_with_moe(stacked)
|
||||||
# Cross Modal Attention
|
|
||||||
attended = self.cross_attention(stacked)
|
|
||||||
|
|
||||||
# Fusion
|
|
||||||
fused = self.fusion(attended)
|
fused = self.fusion(attended)
|
||||||
|
|
||||||
return fused
|
return fused
|
||||||
|
|
||||||
def forward_delivery(
|
def forward_delivery(
|
||||||
@ -315,17 +356,24 @@ class LNPModel(nn.Module):
|
|||||||
if self.mpnn_encoder is not None:
|
if self.mpnn_encoder is not None:
|
||||||
self.mpnn_encoder.clear_cache()
|
self.mpnn_encoder.clear_cache()
|
||||||
|
|
||||||
|
def get_last_moe_extras(self) -> Optional[Dict[str, torch.Tensor]]:
|
||||||
|
"""返回最近一次 forward 中 MoE 模块的副产物(aux loss、gates、probs)。
|
||||||
|
|
||||||
|
若未启用 MoE 或还未调用过 forward,返回 None。
|
||||||
|
"""
|
||||||
|
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(不含任务头)。
|
||||||
|
|
||||||
包含: token_projector, cross_attention, fusion
|
包含: token_projector, cross_attention, fusion,以及(启用时)moe。
|
||||||
"""
|
"""
|
||||||
backbone_keys = []
|
backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.")
|
||||||
for name in self.state_dict().keys():
|
backbone_keys = [
|
||||||
if name.startswith(("token_projector.", "cross_attention.", "fusion.")):
|
name for name in self.state_dict().keys()
|
||||||
backbone_keys.append(name)
|
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 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]:
|
||||||
@ -343,25 +391,25 @@ class LNPModel(nn.Module):
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
从预训练 checkpoint 加载 backbone 和(可选)delivery head 权重。
|
从预训练 checkpoint 加载 backbone 和(可选)delivery head 权重。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pretrain_state_dict: 预训练模型的 state_dict
|
pretrain_state_dict: 预训练模型的 state_dict
|
||||||
load_delivery_head: 是否加载 delivery head 权重
|
load_delivery_head: 是否加载 delivery head 权重
|
||||||
strict: 是否严格匹配(默认 False,允许缺失/多余的键)
|
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():
|
||||||
# Backbone 部分
|
if name.startswith(backbone_prefixes):
|
||||||
if name.startswith(("token_projector.", "cross_attention.", "fusion.")):
|
|
||||||
keys_to_load.append(name)
|
keys_to_load.append(name)
|
||||||
# Delivery head(可选)
|
|
||||||
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)
|
||||||
|
|
||||||
filtered_state_dict = {k: v for k, v in pretrain_state_dict.items() if k in keys_to_load}
|
filtered_state_dict = {
|
||||||
|
k: v for k, v in pretrain_state_dict.items() if k in keys_to_load
|
||||||
# 加载权重
|
}
|
||||||
|
|
||||||
missing, unexpected = [], []
|
missing, 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():
|
||||||
@ -369,12 +417,14 @@ class LNPModel(nn.Module):
|
|||||||
if model_state[k].shape == v.shape:
|
if model_state[k].shape == v.shape:
|
||||||
model_state[k] = v
|
model_state[k] = v
|
||||||
else:
|
else:
|
||||||
unexpected.append(f"{k} (shape mismatch: {model_state[k].shape} vs {v.shape})")
|
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 (missing or unexpected):
|
||||||
raise RuntimeError(f"Missing keys: {missing}, Unexpected keys: {unexpected}")
|
raise RuntimeError(f"Missing keys: {missing}, Unexpected keys: {unexpected}")
|
||||||
|
|
||||||
@ -391,6 +441,12 @@ class LNPModelWithoutMPNN(LNPModel):
|
|||||||
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 相关(新增) ============
|
||||||
|
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,
|
||||||
) -> None:
|
) -> None:
|
||||||
# 移除 mpnn 维度
|
# 移除 mpnn 维度
|
||||||
dims = input_dims or DEFAULT_INPUT_DIMS.copy()
|
dims = input_dims or DEFAULT_INPUT_DIMS.copy()
|
||||||
@ -406,5 +462,10 @@ class LNPModelWithoutMPNN(LNPModel):
|
|||||||
mpnn_checkpoint=None,
|
mpnn_checkpoint=None,
|
||||||
mpnn_ensemble_paths=None,
|
mpnn_ensemble_paths=None,
|
||||||
input_dims=dims,
|
input_dims=dims,
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -178,8 +178,21 @@ def create_model(
|
|||||||
dropout: float = 0.1,
|
dropout: float = 0.1,
|
||||||
use_mpnn: bool = False,
|
use_mpnn: bool = False,
|
||||||
mpnn_device: str = "cpu",
|
mpnn_device: str = "cpu",
|
||||||
|
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,
|
||||||
) -> Union[LNPModel, LNPModelWithoutMPNN]:
|
) -> Union[LNPModel, LNPModelWithoutMPNN]:
|
||||||
"""创建模型"""
|
"""创建模型"""
|
||||||
|
moe_kwargs = dict(
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
if use_mpnn:
|
if use_mpnn:
|
||||||
ensemble_paths = find_mpnn_ensemble_paths()
|
ensemble_paths = find_mpnn_ensemble_paths()
|
||||||
return LNPModel(
|
return LNPModel(
|
||||||
@ -191,6 +204,7 @@ def create_model(
|
|||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
mpnn_ensemble_paths=ensemble_paths,
|
mpnn_ensemble_paths=ensemble_paths,
|
||||||
mpnn_device=mpnn_device,
|
mpnn_device=mpnn_device,
|
||||||
|
**moe_kwargs,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return LNPModelWithoutMPNN(
|
return LNPModelWithoutMPNN(
|
||||||
@ -200,6 +214,7 @@ def create_model(
|
|||||||
fusion_strategy=fusion_strategy,
|
fusion_strategy=fusion_strategy,
|
||||||
head_hidden_dim=head_hidden_dim,
|
head_hidden_dim=head_hidden_dim,
|
||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
|
**moe_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@ -344,7 +359,6 @@ def load_pretrain_weights_to_model(
|
|||||||
|
|
||||||
|
|
||||||
# ============ 内层 Optuna 调参 ============
|
# ============ 内层 Optuna 调参 ============
|
||||||
|
|
||||||
def run_inner_optuna(
|
def run_inner_optuna(
|
||||||
full_dataset: LNPDataset,
|
full_dataset: LNPDataset,
|
||||||
inner_train_indices: np.ndarray,
|
inner_train_indices: np.ndarray,
|
||||||
@ -362,6 +376,11 @@ def run_inner_optuna(
|
|||||||
pretrain_config: Optional[Dict] = None,
|
pretrain_config: Optional[Dict] = None,
|
||||||
load_delivery_head: bool = True,
|
load_delivery_head: bool = True,
|
||||||
rdkit_cache: Optional[Dict] = None,
|
rdkit_cache: Optional[Dict] = None,
|
||||||
|
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,
|
||||||
) -> Tuple[Dict, int, optuna.Study]:
|
) -> Tuple[Dict, int, optuna.Study]:
|
||||||
"""
|
"""
|
||||||
在内层数据上运行 Optuna 超参搜索。
|
在内层数据上运行 Optuna 超参搜索。
|
||||||
@ -456,6 +475,11 @@ def run_inner_optuna(
|
|||||||
dropout=dropout,
|
dropout=dropout,
|
||||||
use_mpnn=use_mpnn,
|
use_mpnn=use_mpnn,
|
||||||
mpnn_device=device.type,
|
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_jitter_noise=moe_jitter_noise,
|
||||||
)
|
)
|
||||||
if rdkit_cache is not None:
|
if rdkit_cache is not None:
|
||||||
model.rdkit_encoder._cache = rdkit_cache
|
model.rdkit_encoder._cache = rdkit_cache
|
||||||
@ -524,7 +548,6 @@ def run_inner_optuna(
|
|||||||
|
|
||||||
|
|
||||||
# ============ 单 fold 执行(可跨进程调用) ============
|
# ============ 单 fold 执行(可跨进程调用) ============
|
||||||
|
|
||||||
def _run_single_outer_fold(
|
def _run_single_outer_fold(
|
||||||
outer_fold: int,
|
outer_fold: int,
|
||||||
outer_train_idx: np.ndarray,
|
outer_train_idx: np.ndarray,
|
||||||
@ -543,6 +566,11 @@ def _run_single_outer_fold(
|
|||||||
pretrain_config: Optional[Dict],
|
pretrain_config: Optional[Dict],
|
||||||
load_delivery_head: bool,
|
load_delivery_head: bool,
|
||||||
device_str: str,
|
device_str: str,
|
||||||
|
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,
|
||||||
) -> Dict:
|
) -> Dict:
|
||||||
"""
|
"""
|
||||||
执行单个外层 fold 的完整流程(内层调参 + 外层训练 + 评估)。
|
执行单个外层 fold 的完整流程(内层调参 + 外层训练 + 评估)。
|
||||||
@ -591,6 +619,11 @@ def _run_single_outer_fold(
|
|||||||
pretrain_config=pretrain_config,
|
pretrain_config=pretrain_config,
|
||||||
load_delivery_head=load_delivery_head,
|
load_delivery_head=load_delivery_head,
|
||||||
rdkit_cache=rdkit_cache,
|
rdkit_cache=rdkit_cache,
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 保存最佳参数
|
# 保存最佳参数
|
||||||
@ -624,6 +657,11 @@ def _run_single_outer_fold(
|
|||||||
dropout=best_params["dropout"],
|
dropout=best_params["dropout"],
|
||||||
use_mpnn=use_mpnn,
|
use_mpnn=use_mpnn,
|
||||||
mpnn_device=device.type,
|
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_jitter_noise=moe_jitter_noise,
|
||||||
)
|
)
|
||||||
model.rdkit_encoder._cache = rdkit_cache
|
model.rdkit_encoder._cache = rdkit_cache
|
||||||
|
|
||||||
@ -659,6 +697,11 @@ def _run_single_outer_fold(
|
|||||||
"head_hidden_dim": best_params["head_hidden_dim"],
|
"head_hidden_dim": best_params["head_hidden_dim"],
|
||||||
"dropout": best_params["dropout"],
|
"dropout": best_params["dropout"],
|
||||||
"use_mpnn": use_mpnn,
|
"use_mpnn": use_mpnn,
|
||||||
|
"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,
|
||||||
}
|
}
|
||||||
|
|
||||||
torch.save({
|
torch.save({
|
||||||
@ -696,7 +739,6 @@ def _run_single_outer_fold(
|
|||||||
|
|
||||||
|
|
||||||
# ============ 主流程 ============
|
# ============ 主流程 ============
|
||||||
|
|
||||||
@app.command()
|
@app.command()
|
||||||
def main(
|
def main(
|
||||||
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
|
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
|
||||||
@ -717,6 +759,12 @@ def main(
|
|||||||
load_delivery_head: bool = False,
|
load_delivery_head: bool = False,
|
||||||
# MPNN
|
# MPNN
|
||||||
use_mpnn: bool = False,
|
use_mpnn: bool = False,
|
||||||
|
# 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,
|
||||||
# 并行
|
# 并行
|
||||||
parallel: bool = False,
|
parallel: bool = False,
|
||||||
# 设备
|
# 设备
|
||||||
@ -805,6 +853,11 @@ def main(
|
|||||||
pretrain_config=pretrain_config,
|
pretrain_config=pretrain_config,
|
||||||
load_delivery_head=load_delivery_head,
|
load_delivery_head=load_delivery_head,
|
||||||
device_str=device_str,
|
device_str=device_str,
|
||||||
|
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,
|
||||||
))
|
))
|
||||||
|
|
||||||
if parallel:
|
if parallel:
|
||||||
|
|||||||
@ -19,6 +19,7 @@ class LossWeights:
|
|||||||
delivery: float = 1.0
|
delivery: float = 1.0
|
||||||
biodist: float = 1.0
|
biodist: float = 1.0
|
||||||
toxic: float = 1.0
|
toxic: float = 1.0
|
||||||
|
moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效)
|
||||||
# size: float = 0.1
|
# size: float = 0.1
|
||||||
# pdi: float = 0.3
|
# pdi: float = 0.3
|
||||||
# ee: float = 0.3
|
# ee: float = 0.3
|
||||||
@ -32,23 +33,26 @@ def compute_multitask_loss(
|
|||||||
targets: Dict[str, torch.Tensor],
|
targets: Dict[str, torch.Tensor],
|
||||||
mask: Dict[str, torch.Tensor],
|
mask: Dict[str, torch.Tensor],
|
||||||
weights: Optional[LossWeights] = None,
|
weights: Optional[LossWeights] = None,
|
||||||
|
model: Optional[nn.Module] = None,
|
||||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||||
"""
|
"""
|
||||||
计算多任务损失。
|
计算多任务损失。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
outputs: 模型输出
|
outputs: 模型输出
|
||||||
targets: 真实标签
|
targets: 真实标签
|
||||||
mask: 有效样本掩码
|
mask: 有效样本掩码
|
||||||
weights: 各任务权重
|
weights: 各任务权重
|
||||||
|
model: 可选,若提供且模型暴露 get_last_moe_extras(),会把 MoE 的
|
||||||
|
load-balancing aux loss 按 weights.moe_lb 加进总 loss。
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(total_loss, loss_dict) 总损失和各任务损失
|
(total_loss, loss_dict) 总损失和各任务损失
|
||||||
"""
|
"""
|
||||||
weights = weights or LossWeights()
|
weights = weights or LossWeights()
|
||||||
losses = {}
|
losses = {}
|
||||||
total_loss = torch.tensor(0.0, device=next(iter(outputs.values())).device)
|
total_loss = torch.tensor(0.0, device=next(iter(outputs.values())).device)
|
||||||
|
|
||||||
# size: MSE loss
|
# size: MSE loss
|
||||||
if "size" in targets and mask["size"].any():
|
if "size" in targets and mask["size"].any():
|
||||||
m = mask["size"]
|
m = mask["size"]
|
||||||
@ -56,7 +60,7 @@ def compute_multitask_loss(
|
|||||||
tgt = targets["size"][m]
|
tgt = targets["size"][m]
|
||||||
losses["size"] = F.mse_loss(pred, tgt)
|
losses["size"] = F.mse_loss(pred, tgt)
|
||||||
total_loss = total_loss + weights.size * losses["size"]
|
total_loss = total_loss + weights.size * losses["size"]
|
||||||
|
|
||||||
# delivery: MSE loss
|
# delivery: MSE loss
|
||||||
if "delivery" in targets and mask["delivery"].any():
|
if "delivery" in targets and mask["delivery"].any():
|
||||||
m = mask["delivery"]
|
m = mask["delivery"]
|
||||||
@ -64,7 +68,7 @@ def compute_multitask_loss(
|
|||||||
tgt = targets["delivery"][m]
|
tgt = targets["delivery"][m]
|
||||||
losses["delivery"] = F.mse_loss(pred, tgt)
|
losses["delivery"] = F.mse_loss(pred, tgt)
|
||||||
total_loss = total_loss + weights.delivery * losses["delivery"]
|
total_loss = total_loss + weights.delivery * losses["delivery"]
|
||||||
|
|
||||||
# pdi: CrossEntropy
|
# pdi: CrossEntropy
|
||||||
if "pdi" in targets and mask["pdi"].any():
|
if "pdi" in targets and mask["pdi"].any():
|
||||||
m = mask["pdi"]
|
m = mask["pdi"]
|
||||||
@ -72,7 +76,7 @@ def compute_multitask_loss(
|
|||||||
tgt = targets["pdi"][m]
|
tgt = targets["pdi"][m]
|
||||||
losses["pdi"] = F.cross_entropy(pred, tgt)
|
losses["pdi"] = F.cross_entropy(pred, tgt)
|
||||||
total_loss = total_loss + weights.pdi * losses["pdi"]
|
total_loss = total_loss + weights.pdi * losses["pdi"]
|
||||||
|
|
||||||
# ee: CrossEntropy
|
# ee: CrossEntropy
|
||||||
if "ee" in targets and mask["ee"].any():
|
if "ee" in targets and mask["ee"].any():
|
||||||
m = mask["ee"]
|
m = mask["ee"]
|
||||||
@ -80,7 +84,7 @@ def compute_multitask_loss(
|
|||||||
tgt = targets["ee"][m]
|
tgt = targets["ee"][m]
|
||||||
losses["ee"] = F.cross_entropy(pred, tgt)
|
losses["ee"] = F.cross_entropy(pred, tgt)
|
||||||
total_loss = total_loss + weights.ee * losses["ee"]
|
total_loss = total_loss + weights.ee * losses["ee"]
|
||||||
|
|
||||||
# toxic: CrossEntropy
|
# toxic: CrossEntropy
|
||||||
if "toxic" in targets and mask["toxic"].any():
|
if "toxic" in targets and mask["toxic"].any():
|
||||||
m = mask["toxic"]
|
m = mask["toxic"]
|
||||||
@ -88,20 +92,26 @@ def compute_multitask_loss(
|
|||||||
tgt = targets["toxic"][m]
|
tgt = targets["toxic"][m]
|
||||||
losses["toxic"] = F.cross_entropy(pred, tgt)
|
losses["toxic"] = F.cross_entropy(pred, tgt)
|
||||||
total_loss = total_loss + weights.toxic * losses["toxic"]
|
total_loss = total_loss + weights.toxic * losses["toxic"]
|
||||||
|
|
||||||
# biodist: KL divergence
|
# biodist: KL divergence
|
||||||
if "biodist" in targets and mask["biodist"].any():
|
if "biodist" in targets and mask["biodist"].any():
|
||||||
m = mask["biodist"]
|
m = mask["biodist"]
|
||||||
pred = outputs["biodist"][m]
|
pred = outputs["biodist"][m]
|
||||||
tgt = targets["biodist"][m]
|
tgt = targets["biodist"][m]
|
||||||
# KL divergence: KL(target || pred)
|
|
||||||
losses["biodist"] = F.kl_div(
|
losses["biodist"] = F.kl_div(
|
||||||
pred.log().clamp(min=-100),
|
pred.log().clamp(min=-100),
|
||||||
tgt,
|
tgt,
|
||||||
reduction="batchmean",
|
reduction="batchmean",
|
||||||
)
|
)
|
||||||
total_loss = total_loss + weights.biodist * losses["biodist"]
|
total_loss = total_loss + weights.biodist * losses["biodist"]
|
||||||
|
|
||||||
|
# MoE load-balancing aux loss(可选)
|
||||||
|
if model is not None and hasattr(model, "get_last_moe_extras"):
|
||||||
|
extras = model.get_last_moe_extras()
|
||||||
|
if extras is not None and "lb_loss" in extras:
|
||||||
|
losses["moe_lb"] = extras["lb_loss"]
|
||||||
|
total_loss = total_loss + weights.moe_lb * losses["moe_lb"]
|
||||||
|
|
||||||
return total_loss, losses
|
return total_loss, losses
|
||||||
|
|
||||||
|
|
||||||
@ -115,7 +125,7 @@ def train_epoch(
|
|||||||
"""训练一个 epoch"""
|
"""训练一个 epoch"""
|
||||||
model.train()
|
model.train()
|
||||||
total_loss = 0.0
|
total_loss = 0.0
|
||||||
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic"]}
|
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic", "moe_lb"]}
|
||||||
n_batches = 0
|
n_batches = 0
|
||||||
|
|
||||||
for batch in tqdm(loader, desc="Training", leave=False):
|
for batch in tqdm(loader, desc="Training", leave=False):
|
||||||
@ -126,7 +136,7 @@ def train_epoch(
|
|||||||
|
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
outputs = model(smiles, tabular)
|
outputs = model(smiles, tabular)
|
||||||
loss, losses = compute_multitask_loss(outputs, targets, mask, weights)
|
loss, losses = compute_multitask_loss(outputs, targets, mask, weights, model=model)
|
||||||
|
|
||||||
loss.backward()
|
loss.backward()
|
||||||
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||||
@ -153,7 +163,7 @@ def validate(
|
|||||||
"""验证"""
|
"""验证"""
|
||||||
model.eval()
|
model.eval()
|
||||||
total_loss = 0.0
|
total_loss = 0.0
|
||||||
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic"]}
|
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic", "moe_lb"]}
|
||||||
n_batches = 0
|
n_batches = 0
|
||||||
|
|
||||||
# 用于计算准确率
|
# 用于计算准确率
|
||||||
@ -167,7 +177,7 @@ def validate(
|
|||||||
mask = {k: v.to(device) for k, v in batch["mask"].items()}
|
mask = {k: v.to(device) for k, v in batch["mask"].items()}
|
||||||
|
|
||||||
outputs = model(smiles, tabular)
|
outputs = model(smiles, tabular)
|
||||||
loss, losses = compute_multitask_loss(outputs, targets, mask, weights)
|
loss, losses = compute_multitask_loss(outputs, targets, mask, weights, model=model)
|
||||||
|
|
||||||
total_loss += loss.item()
|
total_loss += loss.item()
|
||||||
for k, v in losses.items():
|
for k, v in losses.items():
|
||||||
|
|||||||
@ -28,6 +28,7 @@ class LossWeightsBalanced:
|
|||||||
delivery: float = 1.0
|
delivery: float = 1.0
|
||||||
biodist: float = 1.0
|
biodist: float = 1.0
|
||||||
toxic: float = 1.0
|
toxic: float = 1.0
|
||||||
|
moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效)
|
||||||
|
|
||||||
|
|
||||||
def compute_class_weights_from_loader(
|
def compute_class_weights_from_loader(
|
||||||
@ -110,6 +111,7 @@ def compute_multitask_loss_balanced(
|
|||||||
mask: Dict[str, torch.Tensor],
|
mask: Dict[str, torch.Tensor],
|
||||||
task_weights: Optional[LossWeightsBalanced] = None,
|
task_weights: Optional[LossWeightsBalanced] = None,
|
||||||
class_weights: Optional[ClassWeights] = None,
|
class_weights: Optional[ClassWeights] = None,
|
||||||
|
model: Optional[nn.Module] = None,
|
||||||
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||||
"""
|
"""
|
||||||
计算带类权重的多任务损失。
|
计算带类权重的多任务损失。
|
||||||
@ -185,6 +187,13 @@ def compute_multitask_loss_balanced(
|
|||||||
reduction="batchmean",
|
reduction="batchmean",
|
||||||
)
|
)
|
||||||
total_loss = total_loss + task_weights.biodist * losses["biodist"]
|
total_loss = total_loss + task_weights.biodist * losses["biodist"]
|
||||||
|
|
||||||
|
# MoE load-balancing aux loss(可选)
|
||||||
|
if model is not None and hasattr(model, "get_last_moe_extras"):
|
||||||
|
extras = model.get_last_moe_extras()
|
||||||
|
if extras is not None and "lb_loss" in extras:
|
||||||
|
losses["moe_lb"] = extras["lb_loss"]
|
||||||
|
total_loss = total_loss + task_weights.moe_lb * losses["moe_lb"]
|
||||||
|
|
||||||
return total_loss, losses
|
return total_loss, losses
|
||||||
|
|
||||||
@ -200,7 +209,7 @@ def train_epoch_balanced(
|
|||||||
"""带类权重的训练一个 epoch"""
|
"""带类权重的训练一个 epoch"""
|
||||||
model.train()
|
model.train()
|
||||||
total_loss = 0.0
|
total_loss = 0.0
|
||||||
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic"]}
|
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic", "moe_lb"]}
|
||||||
n_batches = 0
|
n_batches = 0
|
||||||
|
|
||||||
for batch in tqdm(loader, desc="Training", leave=False):
|
for batch in tqdm(loader, desc="Training", leave=False):
|
||||||
@ -212,7 +221,7 @@ def train_epoch_balanced(
|
|||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
outputs = model(smiles, tabular)
|
outputs = model(smiles, tabular)
|
||||||
loss, losses = compute_multitask_loss_balanced(
|
loss, losses = compute_multitask_loss_balanced(
|
||||||
outputs, targets, mask, task_weights, class_weights
|
outputs, targets, mask, task_weights, class_weights, model=model,
|
||||||
)
|
)
|
||||||
|
|
||||||
loss.backward()
|
loss.backward()
|
||||||
@ -241,7 +250,7 @@ def validate_balanced(
|
|||||||
"""带类权重的验证"""
|
"""带类权重的验证"""
|
||||||
model.eval()
|
model.eval()
|
||||||
total_loss = 0.0
|
total_loss = 0.0
|
||||||
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic"]}
|
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic", "moe_lb"]}
|
||||||
n_batches = 0
|
n_batches = 0
|
||||||
|
|
||||||
# 用于计算准确率
|
# 用于计算准确率
|
||||||
@ -256,7 +265,7 @@ def validate_balanced(
|
|||||||
|
|
||||||
outputs = model(smiles, tabular)
|
outputs = model(smiles, tabular)
|
||||||
loss, losses = compute_multitask_loss_balanced(
|
loss, losses = compute_multitask_loss_balanced(
|
||||||
outputs, targets, mask, task_weights, class_weights
|
outputs, targets, mask, task_weights, class_weights, model=model,
|
||||||
)
|
)
|
||||||
|
|
||||||
total_loss += loss.item()
|
total_loss += loss.item()
|
||||||
@ -286,7 +295,7 @@ def validate_balanced(
|
|||||||
return metrics
|
return metrics
|
||||||
|
|
||||||
|
|
||||||
BACKBONE_PREFIXES = ("token_projector.", "cross_attention.", "fusion.")
|
BACKBONE_PREFIXES = ("token_projector.", "cross_attention.", "fusion.", "moe.")
|
||||||
|
|
||||||
|
|
||||||
def build_optimizer(
|
def build_optimizer(
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user