Add Mixture-of-Experts (MoE) layer support across modeling pipeline

This commit is contained in:
Michelle0474 2026-05-26 15:26:18 +08:00
parent 804bda5576
commit cd4534f05a
9 changed files with 491 additions and 69 deletions

View File

@ -18,6 +18,14 @@ USE_SWA_FLAG = $(if $(USE_SWA),--use-swa,)
PARALLEL_FLAG = $(if $(PARALLEL),--parallel,)
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 #
#################################################################################
@ -117,10 +125,12 @@ pretrain: requirements
train: requirements
$(PYTHON_INTERPRETER) -m lnp_ml.modeling.nested_cv_optuna \
$(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 \
$(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) #

View File

@ -41,8 +41,12 @@ BIODIST_ORGANS = ["lymph_nodes", "heart", "liver", "spleen", "lung", "kidney", "
def get_token_names(model: Union[LNPModel, LNPModelWithoutMPNN]) -> List[str]:
if model.use_mpnn:
return ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
return ["morgan", "maccs", "desc", "comp", "phys", "help", "exp"]
names = ["mpnn", "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] = []
for start in range(0, all_tokens.size(0), batch_size):
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)
all_weights.append(weights.cpu())
@ -295,7 +300,7 @@ def plot_token_importance(
vals_sorted = normed[order]
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])
colors = [color_a if n in channel_a_set else color_b for n in names_sorted]

View File

@ -176,8 +176,22 @@ def create_model(
dropout: float = 0.1,
use_mpnn: bool = False,
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]:
"""创建模型"""
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:
ensemble_paths = find_mpnn_ensemble_paths()
return LNPModel(
@ -189,6 +203,7 @@ def create_model(
dropout=dropout,
mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device,
**moe_kwargs,
)
else:
return LNPModelWithoutMPNN(
@ -198,6 +213,7 @@ def create_model(
fusion_strategy=fusion_strategy,
head_hidden_dim=head_hidden_dim,
dropout=dropout,
**moe_kwargs,
)
@ -232,7 +248,6 @@ def load_pretrain_weights_to_model(
# ============ 3-fold Optuna 调参 ============
def run_optuna_cv(
full_dataset: LNPDataset,
strata: np.ndarray,
@ -249,6 +264,11 @@ def run_optuna_cv(
pretrain_config: Optional[Dict] = None,
load_delivery_head: bool = True,
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]:
"""
使用全量数据做 3-fold CV Optuna 超参搜索
@ -335,6 +355,11 @@ def run_optuna_cv(
dropout=dropout,
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_jitter_noise=moe_jitter_noise,
)
if rdkit_cache is not None:
model.rdkit_encoder._cache = rdkit_cache
@ -404,7 +429,6 @@ def run_optuna_cv(
# ============ 主流程 ============
@app.command()
def main(
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
@ -427,6 +451,12 @@ def main(
load_delivery_head: bool = False,
# MPNN
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",
):
@ -507,6 +537,11 @@ def main(
pretrain_config=pretrain_config,
load_delivery_head=load_delivery_head,
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"],
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_jitter_noise=moe_jitter_noise,
)
model.rdkit_encoder._cache = rdkit_cache
@ -611,6 +651,11 @@ def main(
"head_hidden_dim": best_params["head_hidden_dim"],
"dropout": best_params["dropout"],
"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({

View File

@ -1,6 +1,6 @@
from lnp_ml.modeling.layers.token_projector import TokenProjector
from lnp_ml.modeling.layers.bidirectional_cross_attention import CrossModalAttention
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"]

View 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 lossgates
"""
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):
"""单个专家 MLPin_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] 化学侧 tokenrouter 不看
tab: [B, T_tab, d_model] 配方/实验侧 tokenrouter
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/Kloss 1
"""
f = expert_mask.mean(dim=0) # [K]
p = probs_full.mean(dim=0) # [K]
return self.n_experts * (f * p).sum()

View File

@ -5,7 +5,7 @@ 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
from lnp_ml.modeling.layers import TokenProjector, CrossModalAttention, FusionLayer, MoEBlock
from lnp_ml.modeling.heads import MultiTaskHead
@ -62,6 +62,12 @@ class LNPModel(nn.Module):
mpnn_device: str = "cpu",
# 输入维度配置
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:
super().__init__()
@ -107,10 +113,28 @@ class LNPModel(nn.Module):
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 ============
self.fusion = FusionLayer(
d_model=d_model,
n_tokens=n_tokens,
n_tokens=n_fusion_tokens,
strategy=fusion_strategy,
)
@ -169,6 +193,29 @@ class LNPModel(nn.Module):
stacked = torch.stack([projected[k] for k in token_order], dim=1)
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_extrastrainer 端通过 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(
self,
stacked: torch.Tensor,
@ -178,14 +225,14 @@ class LNPModel(nn.Module):
从已投影的 stacked tokens 开始 forward用于 Captum 归因
Args:
stacked: [B, n_tokens, d_model] TokenProjector 输出后 stack 的张量
stacked: [B, n_tokens, d_model] TokenProjector 输出后 stack 的张量
task: 指定单任务名 ("size", "pdi", "ee", "delivery", "biodist", "toxic")
若为 None返回 delivery head 的标量输出
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)
if task is None:
@ -240,7 +287,7 @@ class LNPModel(nn.Module):
tabular: Dict[str, torch.Tensor],
) -> torch.Tensor:
"""
Backbone forward编码 -> 投影 -> 注意力 -> 融合不经过任务头
Backbone forward编码 -> 投影 -> 注意力 -> (可选 MoE) -> 融合不经过任务头
用于 pretrain 阶段或需要提取特征的场景
@ -251,15 +298,9 @@ class LNPModel(nn.Module):
Returns:
fused: [B, fusion_dim] 融合后的特征向量
"""
# 编码 + 投影 + stack
stacked = self._encode_and_project(smiles, tabular)
# Cross Modal Attention
attended = self.cross_attention(stacked)
# Fusion
attended = self._attended_with_moe(stacked)
fused = self.fusion(attended)
return fused
def forward_delivery(
@ -315,17 +356,24 @@ class LNPModel(nn.Module):
if self.mpnn_encoder is not None:
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]:
"""
获取 backbone 部分的 state_dict不含任务头
包含: token_projector, cross_attention, fusion
包含: token_projector, cross_attention, fusion以及启用时moe
"""
backbone_keys = []
for name in self.state_dict().keys():
if name.startswith(("token_projector.", "cross_attention.", "fusion.")):
backbone_keys.append(name)
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}
def get_delivery_head_state_dict(self) -> Dict[str, torch.Tensor]:
@ -349,19 +397,19 @@ class LNPModel(nn.Module):
load_delivery_head: 是否加载 delivery head 权重
strict: 是否严格匹配默认 False允许缺失/多余的键
"""
# 筛选要加载的参数
backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.")
keys_to_load = []
for name in pretrain_state_dict.keys():
# Backbone 部分
if name.startswith(("token_projector.", "cross_attention.", "fusion.")):
if name.startswith(backbone_prefixes):
keys_to_load.append(name)
# Delivery head可选
elif load_delivery_head and name.startswith("head.delivery_head."):
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 = [], []
model_state = self.state_dict()
for k, v in filtered_state_dict.items():
@ -369,7 +417,9 @@ class LNPModel(nn.Module):
if model_state[k].shape == v.shape:
model_state[k] = v
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:
unexpected.append(k)
@ -391,6 +441,12 @@ class LNPModelWithoutMPNN(LNPModel):
head_hidden_dim: int = 128,
dropout: float = 0.1,
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:
# 移除 mpnn 维度
dims = input_dims or DEFAULT_INPUT_DIMS.copy()
@ -406,5 +462,10 @@ class LNPModelWithoutMPNN(LNPModel):
mpnn_checkpoint=None,
mpnn_ensemble_paths=None,
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,
)

View File

@ -178,8 +178,21 @@ def create_model(
dropout: float = 0.1,
use_mpnn: bool = False,
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]:
"""创建模型"""
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:
ensemble_paths = find_mpnn_ensemble_paths()
return LNPModel(
@ -191,6 +204,7 @@ def create_model(
dropout=dropout,
mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device,
**moe_kwargs,
)
else:
return LNPModelWithoutMPNN(
@ -200,6 +214,7 @@ def create_model(
fusion_strategy=fusion_strategy,
head_hidden_dim=head_hidden_dim,
dropout=dropout,
**moe_kwargs,
)
@ -344,7 +359,6 @@ def load_pretrain_weights_to_model(
# ============ 内层 Optuna 调参 ============
def run_inner_optuna(
full_dataset: LNPDataset,
inner_train_indices: np.ndarray,
@ -362,6 +376,11 @@ def run_inner_optuna(
pretrain_config: Optional[Dict] = None,
load_delivery_head: bool = True,
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]:
"""
在内层数据上运行 Optuna 超参搜索
@ -456,6 +475,11 @@ def run_inner_optuna(
dropout=dropout,
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_jitter_noise=moe_jitter_noise,
)
if rdkit_cache is not None:
model.rdkit_encoder._cache = rdkit_cache
@ -524,7 +548,6 @@ def run_inner_optuna(
# ============ 单 fold 执行(可跨进程调用) ============
def _run_single_outer_fold(
outer_fold: int,
outer_train_idx: np.ndarray,
@ -543,6 +566,11 @@ def _run_single_outer_fold(
pretrain_config: Optional[Dict],
load_delivery_head: bool,
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:
"""
执行单个外层 fold 的完整流程内层调参 + 外层训练 + 评估
@ -591,6 +619,11 @@ def _run_single_outer_fold(
pretrain_config=pretrain_config,
load_delivery_head=load_delivery_head,
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"],
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_jitter_noise=moe_jitter_noise,
)
model.rdkit_encoder._cache = rdkit_cache
@ -659,6 +697,11 @@ def _run_single_outer_fold(
"head_hidden_dim": best_params["head_hidden_dim"],
"dropout": best_params["dropout"],
"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({
@ -696,7 +739,6 @@ def _run_single_outer_fold(
# ============ 主流程 ============
@app.command()
def main(
input_path: Path = INTERIM_DATA_DIR / "internal.csv",
@ -717,6 +759,12 @@ def main(
load_delivery_head: bool = False,
# MPNN
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,
# 设备
@ -805,6 +853,11 @@ def main(
pretrain_config=pretrain_config,
load_delivery_head=load_delivery_head,
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:

View File

@ -19,6 +19,7 @@ class LossWeights:
delivery: float = 1.0
biodist: float = 1.0
toxic: float = 1.0
moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效)
# size: float = 0.1
# pdi: float = 0.3
# ee: float = 0.3
@ -32,6 +33,7 @@ def compute_multitask_loss(
targets: Dict[str, torch.Tensor],
mask: Dict[str, torch.Tensor],
weights: Optional[LossWeights] = None,
model: Optional[nn.Module] = None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
"""
计算多任务损失
@ -41,6 +43,8 @@ def compute_multitask_loss(
targets: 真实标签
mask: 有效样本掩码
weights: 各任务权重
model: 可选若提供且模型暴露 get_last_moe_extras()会把 MoE
load-balancing aux loss weights.moe_lb 加进总 loss
Returns:
(total_loss, loss_dict) 总损失和各任务损失
@ -94,7 +98,6 @@ def compute_multitask_loss(
m = mask["biodist"]
pred = outputs["biodist"][m]
tgt = targets["biodist"][m]
# KL divergence: KL(target || pred)
losses["biodist"] = F.kl_div(
pred.log().clamp(min=-100),
tgt,
@ -102,6 +105,13 @@ def compute_multitask_loss(
)
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
@ -115,7 +125,7 @@ def train_epoch(
"""训练一个 epoch"""
model.train()
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
for batch in tqdm(loader, desc="Training", leave=False):
@ -126,7 +136,7 @@ def train_epoch(
optimizer.zero_grad()
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()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
@ -153,7 +163,7 @@ def validate(
"""验证"""
model.eval()
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
# 用于计算准确率
@ -167,7 +177,7 @@ def validate(
mask = {k: v.to(device) for k, v in batch["mask"].items()}
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()
for k, v in losses.items():

View File

@ -28,6 +28,7 @@ class LossWeightsBalanced:
delivery: float = 1.0
biodist: float = 1.0
toxic: float = 1.0
moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效)
def compute_class_weights_from_loader(
@ -110,6 +111,7 @@ def compute_multitask_loss_balanced(
mask: Dict[str, torch.Tensor],
task_weights: Optional[LossWeightsBalanced] = None,
class_weights: Optional[ClassWeights] = None,
model: Optional[nn.Module] = None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
"""
计算带类权重的多任务损失
@ -186,6 +188,13 @@ def compute_multitask_loss_balanced(
)
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
@ -200,7 +209,7 @@ def train_epoch_balanced(
"""带类权重的训练一个 epoch"""
model.train()
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
for batch in tqdm(loader, desc="Training", leave=False):
@ -212,7 +221,7 @@ def train_epoch_balanced(
optimizer.zero_grad()
outputs = model(smiles, tabular)
loss, losses = compute_multitask_loss_balanced(
outputs, targets, mask, task_weights, class_weights
outputs, targets, mask, task_weights, class_weights, model=model,
)
loss.backward()
@ -241,7 +250,7 @@ def validate_balanced(
"""带类权重的验证"""
model.eval()
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
# 用于计算准确率
@ -256,7 +265,7 @@ def validate_balanced(
outputs = model(smiles, tabular)
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()
@ -286,7 +295,7 @@ def validate_balanced(
return metrics
BACKBONE_PREFIXES = ("token_projector.", "cross_attention.", "fusion.")
BACKBONE_PREFIXES = ("token_projector.", "cross_attention.", "fusion.", "moe.")
def build_optimizer(