diff --git a/Makefile b/Makefile index 0664e7a..f86e605 100644 --- a/Makefile +++ b/Makefile @@ -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,11 +125,13 @@ 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) # ################################################################################# diff --git a/lnp_ml/interpretability/token_importance.py b/lnp_ml/interpretability/token_importance.py index 0c851b4..ee0d1d8 100644 --- a/lnp_ml/interpretability/token_importance.py +++ b/lnp_ml/interpretability/token_importance.py @@ -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] diff --git a/lnp_ml/modeling/final_train_optuna_cv.py b/lnp_ml/modeling/final_train_optuna_cv.py index 9e2173f..d3c7a97 100644 --- a/lnp_ml/modeling/final_train_optuna_cv.py +++ b/lnp_ml/modeling/final_train_optuna_cv.py @@ -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,8 +651,13 @@ 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({ "model_state_dict": train_result["final_state"], "config": config, diff --git a/lnp_ml/modeling/layers/__init__.py b/lnp_ml/modeling/layers/__init__.py index f019b23..6fc8fac 100644 --- a/lnp_ml/modeling/layers/__init__.py +++ b/lnp_ml/modeling/layers/__init__.py @@ -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"] \ No newline at end of file diff --git a/lnp_ml/modeling/layers/moe.py b/lnp_ml/modeling/layers/moe.py new file mode 100644 index 0000000..8a07c05 --- /dev/null +++ b/lnp_ml/modeling/layers/moe.py @@ -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() \ No newline at end of file diff --git a/lnp_ml/modeling/models.py b/lnp_ml/modeling/models.py index c85a609..e680a52 100644 --- a/lnp_ml/modeling/models.py +++ b/lnp_ml/modeling/models.py @@ -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_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( 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,26 +287,20 @@ class LNPModel(nn.Module): tabular: Dict[str, torch.Tensor], ) -> torch.Tensor: """ - Backbone forward:编码 -> 投影 -> 注意力 -> 融合,不经过任务头。 - + Backbone forward:编码 -> 投影 -> 注意力 -> (可选 MoE) -> 融合,不经过任务头。 + 用于 pretrain 阶段或需要提取特征的场景。 - + Args: smiles: SMILES 字符串列表,长度为 B tabular: Dict[str, Tensor] - + 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]: @@ -343,25 +391,25 @@ class LNPModel(nn.Module): ) -> None: """ 从预训练 checkpoint 加载 backbone 和(可选)delivery head 权重。 - + Args: pretrain_state_dict: 预训练模型的 state_dict 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,12 +417,14 @@ 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) - + self.load_state_dict(model_state, strict=False) - + if strict and (missing or unexpected): raise RuntimeError(f"Missing keys: {missing}, Unexpected keys: {unexpected}") @@ -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, ) diff --git a/lnp_ml/modeling/nested_cv_optuna.py b/lnp_ml/modeling/nested_cv_optuna.py index 8cc26d0..0fc546a 100644 --- a/lnp_ml/modeling/nested_cv_optuna.py +++ b/lnp_ml/modeling/nested_cv_optuna.py @@ -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: diff --git a/lnp_ml/modeling/trainer.py b/lnp_ml/modeling/trainer.py index dbcd50c..f8c8aed 100644 --- a/lnp_ml/modeling/trainer.py +++ b/lnp_ml/modeling/trainer.py @@ -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,23 +33,26 @@ 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]]: """ 计算多任务损失。 - + Args: outputs: 模型输出 targets: 真实标签 mask: 有效样本掩码 weights: 各任务权重 - + model: 可选,若提供且模型暴露 get_last_moe_extras(),会把 MoE 的 + load-balancing aux loss 按 weights.moe_lb 加进总 loss。 + Returns: (total_loss, loss_dict) 总损失和各任务损失 """ weights = weights or LossWeights() losses = {} total_loss = torch.tensor(0.0, device=next(iter(outputs.values())).device) - + # size: MSE loss if "size" in targets and mask["size"].any(): m = mask["size"] @@ -56,7 +60,7 @@ def compute_multitask_loss( tgt = targets["size"][m] losses["size"] = F.mse_loss(pred, tgt) total_loss = total_loss + weights.size * losses["size"] - + # delivery: MSE loss if "delivery" in targets and mask["delivery"].any(): m = mask["delivery"] @@ -64,7 +68,7 @@ def compute_multitask_loss( tgt = targets["delivery"][m] losses["delivery"] = F.mse_loss(pred, tgt) total_loss = total_loss + weights.delivery * losses["delivery"] - + # pdi: CrossEntropy if "pdi" in targets and mask["pdi"].any(): m = mask["pdi"] @@ -72,7 +76,7 @@ def compute_multitask_loss( tgt = targets["pdi"][m] losses["pdi"] = F.cross_entropy(pred, tgt) total_loss = total_loss + weights.pdi * losses["pdi"] - + # ee: CrossEntropy if "ee" in targets and mask["ee"].any(): m = mask["ee"] @@ -80,7 +84,7 @@ def compute_multitask_loss( tgt = targets["ee"][m] losses["ee"] = F.cross_entropy(pred, tgt) total_loss = total_loss + weights.ee * losses["ee"] - + # toxic: CrossEntropy if "toxic" in targets and mask["toxic"].any(): m = mask["toxic"] @@ -88,20 +92,26 @@ def compute_multitask_loss( tgt = targets["toxic"][m] losses["toxic"] = F.cross_entropy(pred, tgt) total_loss = total_loss + weights.toxic * losses["toxic"] - + # biodist: KL divergence if "biodist" in targets and mask["biodist"].any(): 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, reduction="batchmean", ) 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(): diff --git a/lnp_ml/modeling/trainer_balanced.py b/lnp_ml/modeling/trainer_balanced.py index e67c577..147cd5f 100644 --- a/lnp_ml/modeling/trainer_balanced.py +++ b/lnp_ml/modeling/trainer_balanced.py @@ -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]]: """ 计算带类权重的多任务损失。 @@ -185,6 +187,13 @@ def compute_multitask_loss_balanced( reduction="batchmean", ) 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(