From 54c734dc165c7c908d37350d02fadb3314a6968a Mon Sep 17 00:00:00 2001 From: Michelle0474 <2170308303@qq.com> Date: Mon, 22 Jun 2026 19:39:24 +0800 Subject: [PATCH] feat: integrate MolT5 LLM encoder and fix desc dim --- lnp_ml/modeling/layers/llm_prompt.py | 104 ++ lnp_ml/modeling/models.py | 304 ++-- lnp_ml/modeling/nested_cv_optuna.py | 1954 ++++++++++++++------------ lnp_ml/modeling/predict.py | 16 +- 4 files changed, 1279 insertions(+), 1099 deletions(-) create mode 100644 lnp_ml/modeling/layers/llm_prompt.py diff --git a/lnp_ml/modeling/layers/llm_prompt.py b/lnp_ml/modeling/layers/llm_prompt.py new file mode 100644 index 0000000..a05f8a4 --- /dev/null +++ b/lnp_ml/modeling/layers/llm_prompt.py @@ -0,0 +1,104 @@ +"""LLM 分子特征分支:用 MolT5 直接编码 SMILES 文本。""" + +import os +from typing import Dict, List, Optional + +import torch +import torch.nn as nn + + +# 权重默认路径,可用环境变量 MOLT5_PATH 覆盖 +DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base") + + +class LLMPromptEncoder(nn.Module): + """用 MolT5 编码 SMILES 文本,输出分子特征 F_llm [B, d_model]。 + + - 冻结 encoder 时:对每个 SMILES 缓存其句向量,避免重复前向,降方差、提速。 + - use_lora=True 时:encoder 可训练(LoRA),不缓存。 + """ + + def __init__( + self, + d_model: int, + n_chem_tokens: int = 4, + n_cond_tokens: int = 4, + model_name_or_path: str = DEFAULT_MOLT5_PATH, + freeze: bool = True, + use_lora: bool = False, + lora_r: int = 8, + lora_alpha: int = 16, + lora_dropout: float = 0.05, + max_length: int = 128, + ) -> None: + super().__init__() + from transformers import AutoTokenizer, T5EncoderModel + + self.use_lora = use_lora + self.max_length = max_length + self.tokenizer = AutoTokenizer.from_pretrained(model_name_or_path) + self.encoder = T5EncoderModel.from_pretrained(model_name_or_path) + self.hidden_size = self.encoder.config.d_model + + if use_lora: + self._apply_lora(lora_r, lora_alpha, lora_dropout) + elif freeze: + for p in self.encoder.parameters(): + p.requires_grad = False + self._frozen = freeze and not use_lora + + self.proj_down = nn.Sequential( + nn.Linear(self.hidden_size, d_model), nn.LayerNorm(d_model) + ) + # 冻结特征缓存:smiles -> [H](CPU) + self._cache: Dict[str, torch.Tensor] = {} + + def _apply_lora(self, r: int, alpha: int, dropout: float) -> None: + from peft import LoraConfig, get_peft_model + + for p in self.encoder.parameters(): + p.requires_grad = False + cfg = LoraConfig( + r=r, lora_alpha=alpha, lora_dropout=dropout, + target_modules=["q", "k", "v", "o"], bias="none", + ) + self.encoder = get_peft_model(self.encoder, cfg) + + def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: + m = mask.unsqueeze(-1).float() # [B, L, 1] + return (last_hidden * m).sum(1) / m.sum(1).clamp(min=1e-6) + + @torch.no_grad() + def _encode_frozen(self, smiles: List[str], device: torch.device) -> torch.Tensor: + missing = [s for s in smiles if s not in self._cache] + if missing: + uniq = list(dict.fromkeys(missing)) + for i in range(0, len(uniq), 256): + chunk = uniq[i:i + 256] + enc = self.tokenizer( + chunk, padding=True, truncation=True, + max_length=self.max_length, return_tensors="pt", + ).to(device) + out = self.encoder(**enc).last_hidden_state + pooled = self._mean_pool(out, enc["attention_mask"]) + for s, v in zip(chunk, pooled): + self._cache[s] = v.cpu() + return torch.stack([self._cache[s] for s in smiles]).to(device) + + def _encode_trainable(self, smiles: List[str], device: torch.device) -> torch.Tensor: + enc = self.tokenizer( + list(smiles), padding=True, truncation=True, + max_length=self.max_length, return_tensors="pt", + ).to(device) + out = self.encoder(**enc).last_hidden_state + return self._mean_pool(out, enc["attention_mask"]) + + def forward(self, smiles: List[str], tab: Optional[torch.Tensor] = None) -> torch.Tensor: + """Args: smiles [B] SMILES 字符串列表。Returns: [B, d_model]。""" + device = self.proj_down[0].weight.device + feat = self._encode_frozen(smiles, device) if self._frozen \ + else self._encode_trainable(smiles, device) + return self.proj_down(feat) + + def clear_cache(self) -> None: + self._cache.clear() \ No newline at end of file diff --git a/lnp_ml/modeling/models.py b/lnp_ml/modeling/models.py index e680a52..d312e5d 100644 --- a/lnp_ml/modeling/models.py +++ b/lnp_ml/modeling/models.py @@ -5,20 +5,27 @@ import torch.nn as nn from typing import Dict, List, Optional, Literal from lnp_ml.modeling.encoders import CachedRDKitEncoder, CachedMPNNEncoder -from lnp_ml.modeling.layers import TokenProjector, CrossModalAttention, FusionLayer, MoEBlock +from lnp_ml.modeling.layers import ( + TokenProjector, + SetTransformer, + ResidualConcatFusion, + MoEBlock, + LLMPromptEncoder, +) +from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH from lnp_ml.modeling.heads import MultiTaskHead -PoolingStrategy = Literal["concat", "avg", "max", "attention"] +PoolingStrategy = Literal["attention", "avg", "max"] -# Token 维度配置(根据 ARCHITECTURE.md) +# Token 维度配置 DEFAULT_INPUT_DIMS = { # Channel A: 化学特征 "mpnn": 600, # D-MPNN embedding "morgan": 1024, # Morgan fingerprint "maccs": 167, # MACCS keys - "desc": 210, # RDKit descriptors + "desc": 217, # RDKit descriptors # Channel B: 配方/实验条件 "comp": 5, # 配方比例 "phys": 12, # 物理参数 one-hot @@ -26,8 +33,21 @@ DEFAULT_INPUT_DIMS = { "exp": 32, # 实验条件 one-hot } -# Token 顺序(前 4 个为 Channel A,后 4 个为 Channel B) -TOKEN_ORDER = ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"] +# 化学 / 配方 token 的键顺序 +CHEM_KEYS_WITH_MPNN = ["mpnn", "morgan", "maccs", "desc"] +CHEM_KEYS_NO_MPNN = ["morgan", "maccs", "desc"] +TAB_KEYS = ["comp", "phys", "help", "exp"] + +# backbone 权重前缀(用于预训练加载与导出) +BACKBONE_PREFIXES = ( + "token_projector.", + "set_transformer.", + "fusion.", + "moe.", + "llm_prompt.", +) +# 冻结的 MolT5 encoder 权重前缀,不纳入 backbone(由本地权重加载,不进 checkpoint) +LLM_FROZEN_PREFIX = "llm_prompt.encoder." class LNPModel(nn.Module): @@ -37,37 +57,47 @@ class LNPModel(nn.Module): 架构流程: 1. Encoders: SMILES -> 化学特征; tabular -> 配方/实验特征 2. TokenProjector: 统一到 d_model - 3. Stack: [B, 8, d_model] - 4. CrossModalAttention: Channel A (化学) <-> Channel B (配方/实验) - 5. FusionLayer: [B, 8, d_model] -> [B, fusion_dim] - 6. MultiTaskHead: 多任务预测 + 3. SetTransformer: 对化学 token 集合做置换等变编码 -> chem' + 4. MoE (可选): router 看 tab,expert 吃 chem' -> F_moe + 5. LLM (可选): chem'/tab 注入 MolT5 prompt -> F_llm + 6. ResidualConcatFusion: 拼接 chem'/tab/F_moe/F_llm -> attention pooling + 7. MultiTaskHead: 多任务预测 """ def __init__( self, # 模型维度 d_model: int = 256, - # Cross attention + # Set Transformer num_heads: int = 8, n_attn_layers: int = 4, + set_transformer_block: str = "sab", # Fusion fusion_strategy: PoolingStrategy = "attention", # Head head_hidden_dim: int = 128, # Dropout dropout: float = 0.1, - # MPNN encoder (可选,如果不用 MPNN 可以设为 None) + # MPNN encoder mpnn_checkpoint: Optional[str] = None, mpnn_ensemble_paths: Optional[List[str]] = None, mpnn_device: str = "cpu", # 输入维度配置 input_dims: Optional[Dict[str, int]] = None, - # ============ MoE 相关(新增) ============ + # ============ MoE 相关 ============ use_moe: bool = False, moe_n_experts: int = 4, moe_top_k: int = 2, moe_expert_hidden_mult: int = 2, moe_jitter_noise: float = 0.0, + # ============ LLM 相关 ============ + use_llm: bool = False, + llm_model_path: str = DEFAULT_MOLT5_PATH, + llm_freeze: bool = True, + llm_use_lora: bool = False, + llm_lora_r: int = 8, + llm_lora_alpha: int = 16, + llm_lora_dropout: float = 0.05, ) -> None: super().__init__() @@ -76,10 +106,7 @@ class LNPModel(nn.Module): self.use_mpnn = mpnn_checkpoint is not None or mpnn_ensemble_paths is not None # ============ Encoders ============ - # RDKit encoder (always used) self.rdkit_encoder = CachedRDKitEncoder() - - # MPNN encoder (optional) if self.use_mpnn: self.mpnn_encoder = CachedMPNNEncoder( checkpoint_path=mpnn_checkpoint, @@ -90,27 +117,28 @@ class LNPModel(nn.Module): self.mpnn_encoder = None # ============ Token Projector ============ - # 根据是否使用 MPNN 调整输入维度 proj_input_dims = {k: v for k, v in self.input_dims.items()} if not self.use_mpnn: proj_input_dims.pop("mpnn", None) - self.token_projector = TokenProjector( input_dims=proj_input_dims, d_model=d_model, dropout=dropout, ) - # ============ Cross Modal Attention ============ - n_tokens = 8 if self.use_mpnn else 7 - split_idx = 4 if self.use_mpnn else 3 # Channel A 的 token 数量 + # token 顺序与化学侧 token 数 + self.chem_keys = CHEM_KEYS_WITH_MPNN if self.use_mpnn else CHEM_KEYS_NO_MPNN + self.tab_keys = TAB_KEYS + self.token_order = self.chem_keys + self.tab_keys + self.split_idx = len(self.chem_keys) - self.cross_attention = CrossModalAttention( + # ============ Set Transformer ============ + self.set_transformer = SetTransformer( d_model=d_model, num_heads=num_heads, n_layers=n_attn_layers, - split_idx=split_idx, dropout=dropout, + block=set_transformer_block, ) # ============ MoE Block (可选) ============ @@ -119,24 +147,35 @@ class LNPModel(nn.Module): if use_moe: self.moe = MoEBlock( d_model=d_model, - n_chem_tokens=split_idx, + n_chem_tokens=self.split_idx, n_experts=moe_n_experts, top_k=moe_top_k, expert_hidden_mult=moe_expert_hidden_mult, dropout=dropout, jitter_noise=moe_jitter_noise, ) - n_fusion_tokens = n_tokens + 1 # 多一个 F_moe token else: self.moe = None - n_fusion_tokens = n_tokens - # ============ Fusion Layer ============ - self.fusion = FusionLayer( - d_model=d_model, - n_tokens=n_fusion_tokens, - strategy=fusion_strategy, - ) + # ============ LLM Prompt (可选) ============ + self.use_llm = use_llm + if use_llm: + self.llm_prompt = LLMPromptEncoder( + d_model=d_model, + n_chem_tokens=self.split_idx, + n_cond_tokens=len(self.tab_keys), + model_name_or_path=llm_model_path, + freeze=llm_freeze, + use_lora=llm_use_lora, + lora_r=llm_lora_r, + lora_alpha=llm_lora_alpha, + lora_dropout=llm_lora_dropout, + ) + else: + self.llm_prompt = None + + # ============ Residual Concat + Fusion ============ + self.fusion = ResidualConcatFusion(d_model=d_model, strategy=fusion_strategy) # ============ Multi-Task Head ============ self.head = MultiTaskHead( @@ -151,70 +190,51 @@ class LNPModel(nn.Module): tabular: Dict[str, torch.Tensor], ) -> torch.Tensor: """ - 内部方法:编码 SMILES 和 tabular,返回 stacked tokens。 - + 编码 SMILES 和 tabular,返回 stacked tokens。 + Returns: - stacked: [B, n_tokens, d_model] + stacked: [B, n_tokens, d_model],顺序为 chem 在前、tab 在后 """ - # 获取目标设备(从 tabular 数据推断) device = tabular["comp"].device - - # 1. Encode SMILES - rdkit_features = self.rdkit_encoder(smiles) # {"morgan", "maccs", "desc"} - # 2. 合并所有特征 + rdkit_features = self.rdkit_encoder(smiles) + all_features: Dict[str, torch.Tensor] = {} - - # MPNN 特征(如果启用) if self.use_mpnn: mpnn_features = self.mpnn_encoder(smiles) all_features["mpnn"] = mpnn_features["mpnn"].to(device) - - # RDKit 特征(移到正确设备) all_features["morgan"] = rdkit_features["morgan"].to(device) all_features["maccs"] = rdkit_features["maccs"].to(device) all_features["desc"] = rdkit_features["desc"].to(device) - - # Tabular 特征(已在正确设备上) all_features["comp"] = tabular["comp"] all_features["phys"] = tabular["phys"] all_features["help"] = tabular["help"] all_features["exp"] = tabular["exp"] - # 3. Token Projector: 统一维度 - projected = self.token_projector(all_features) # Dict[str, [B, d_model]] - - # 4. Stack tokens: [B, n_tokens, d_model] - if self.use_mpnn: - token_order = ["mpnn", "morgan", "maccs", "desc", "comp", "phys", "help", "exp"] - else: - token_order = ["morgan", "maccs", "desc", "comp", "phys", "help", "exp"] - - stacked = torch.stack([projected[k] for k in token_order], dim=1) + projected = self.token_projector(all_features) + stacked = torch.stack([projected[k] for k in self.token_order], dim=1) return stacked - def _attended_with_moe(self, stacked: torch.Tensor) -> torch.Tensor: - """ - Cross-attention + 可选 MoE 旁路 → fusion 输入序列。 + def _backbone_from_stacked( + self, stacked: torch.Tensor, smiles: Optional[List[str]] = None + ) -> torch.Tensor: + chem = stacked[:, : self.split_idx, :] + tab = stacked[:, self.split_idx :, :] - - 不启用 MoE 时:返回 [B, n_tokens, d],与原行为一致。 - - 启用 MoE 时:在最后追加一个 F_moe token,返回 [B, n_tokens + 1, d]。 + chem = self.set_transformer(chem) # chem' - 副作用: - 把本次 forward 的 MoE 副产物(aux loss / gates / probs)写到 - self._last_moe_extras,trainer 端通过 get_last_moe_extras() 读取。 - """ - attended = self.cross_attention(stacked) + f_moe = None if self.moe is not None: - split = self.cross_attention.split_idx - chem_prime = attended[:, :split, :] - tab_prime = attended[:, split:, :] - F_moe, extras = self.moe(chem_prime, tab_prime) + f_moe, extras = self.moe(chem, tab) self._last_moe_extras = extras - attended = torch.cat([attended, F_moe.unsqueeze(1)], dim=1) else: self._last_moe_extras = None - return attended + + f_llm = None + if self.llm_prompt is not None and smiles is not None: + f_llm = self.llm_prompt(smiles) + + return self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm) def forward_from_projected( self, @@ -225,15 +245,13 @@ class LNPModel(nn.Module): 从已投影的 stacked tokens 开始 forward,用于 Captum 归因。 Args: - stacked: [B, n_tokens, d_model] TokenProjector 输出后 stack 的张量。 - task: 指定单任务名 ("size", "pdi", "ee", "delivery", "biodist", "toxic")。 - 若为 None,返回 delivery head 的标量输出。 + stacked: [B, n_tokens, d_model] + task: 单任务名;None 时返回 delivery head 输出。 Returns: - [B, 1] 或 [B, num_classes] 对应任务的预测输出。 + 对应任务的预测输出。 """ - attended = self._attended_with_moe(stacked) - fused = self.fusion(attended) + fused = self._backbone_from_stacked(stacked) if task is None: task = "delivery" @@ -259,19 +277,10 @@ class LNPModel(nn.Module): 用原始特征替换 base_projected 中指定 token 的投影,然后 forward。 用于对单个 token 内部特征做 Captum 归因(如 desc 的 210 维)。 - - Args: - raw_feature: [B, input_dim] 某个 token 的原始特征 - feature_key: token 名称,如 "desc" - base_projected: [B, n_tokens, d_model] 其他 token 已投影好的张量 - task: 任务名 - - Returns: - 对应任务的预测输出 """ projected = self.token_projector.projectors[feature_key](raw_feature) gate = torch.sigmoid(self.token_projector.weights[feature_key]) - projected = projected * gate # [B, d_model] + projected = projected * gate token_order = list(self.token_projector.keys) token_idx = token_order.index(feature_key) @@ -286,38 +295,16 @@ class LNPModel(nn.Module): smiles: List[str], tabular: Dict[str, torch.Tensor], ) -> torch.Tensor: - """ - Backbone forward:编码 -> 投影 -> 注意力 -> (可选 MoE) -> 融合,不经过任务头。 - - 用于 pretrain 阶段或需要提取特征的场景。 - - Args: - smiles: SMILES 字符串列表,长度为 B - tabular: Dict[str, Tensor] - - Returns: - fused: [B, fusion_dim] 融合后的特征向量 - """ + """Backbone forward:编码 -> 投影 -> set transformer -> (MoE/LLM) -> 融合。""" stacked = self._encode_and_project(smiles, tabular) - attended = self._attended_with_moe(stacked) - fused = self.fusion(attended) - return fused + return self._backbone_from_stacked(stacked, smiles=smiles) def forward_delivery( self, smiles: List[str], tabular: Dict[str, torch.Tensor], ) -> torch.Tensor: - """ - 仅预测 delivery(用于 pretrain)。 - - Args: - smiles: SMILES 字符串列表,长度为 B - tabular: Dict[str, Tensor] - - Returns: - delivery: [B, 1] 预测的 delivery 值 - """ + """仅预测 delivery(用于 pretrain)。返回 [B, 1]。""" fused = self.forward_backbone(smiles, tabular) return self.head.delivery_head(fused) @@ -328,53 +315,34 @@ class LNPModel(nn.Module): ) -> Dict[str, torch.Tensor]: """ 完整的多任务 forward。 - - Args: - smiles: SMILES 字符串列表,长度为 B - tabular: Dict[str, Tensor],包含: - - "comp": [B, 5] 配方比例 - - "phys": [B, 12] 物理参数 - - "help": [B, 4] Helper lipid - - "exp": [B, 32] 实验条件 Returns: - Dict[str, Tensor]: - - "size": [B, 1] - - "pdi": [B, 4] - - "ee": [B, 3] - - "delivery": [B, 1] - - "biodist": [B, 7] - - "toxic": [B, 2] + Dict[str, Tensor]: size [B,1], pdi [B,4], ee [B,3], + delivery [B,1], biodist [B,7], toxic [B,2] """ fused = self.forward_backbone(smiles, tabular) - outputs = self.head(fused) - return outputs + return self.head(fused) def clear_cache(self) -> None: """清空所有 encoder 的缓存""" self.rdkit_encoder.clear_cache() if self.mpnn_encoder is not None: self.mpnn_encoder.clear_cache() + if self.llm_prompt is not None and hasattr(self.llm_prompt, "clear_cache"): + self.llm_prompt.clear_cache() def get_last_moe_extras(self) -> Optional[Dict[str, torch.Tensor]]: - """返回最近一次 forward 中 MoE 模块的副产物(aux loss、gates、probs)。 - - 若未启用 MoE 或还未调用过 forward,返回 None。 - """ + """返回最近一次 forward 中 MoE 模块的副产物(aux loss、gates、probs)。""" return self._last_moe_extras def get_backbone_state_dict(self) -> Dict[str, torch.Tensor]: """ - 获取 backbone 部分的 state_dict(不含任务头)。 - - 包含: token_projector, cross_attention, fusion,以及(启用时)moe。 + 获取 backbone 部分的 state_dict(不含任务头,且排除冻结的 MolT5 encoder)。 """ - backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.") - backbone_keys = [ - name for name in self.state_dict().keys() - if name.startswith(backbone_prefixes) - ] - return {k: v for k, v in self.state_dict().items() if k in backbone_keys} + return { + k: v for k, v in self.state_dict().items() + if k.startswith(BACKBONE_PREFIXES) and not k.startswith(LLM_FROZEN_PREFIX) + } def get_delivery_head_state_dict(self) -> Dict[str, torch.Tensor]: """获取 delivery head 的 state_dict""" @@ -392,16 +360,11 @@ class LNPModel(nn.Module): """ 从预训练 checkpoint 加载 backbone 和(可选)delivery head 权重。 - Args: - pretrain_state_dict: 预训练模型的 state_dict - load_delivery_head: 是否加载 delivery head 权重 - strict: 是否严格匹配(默认 False,允许缺失/多余的键) + 冻结的 MolT5 encoder 权重不在加载范围内。 """ - backbone_prefixes = ("token_projector.", "cross_attention.", "fusion.", "moe.") - keys_to_load = [] for name in pretrain_state_dict.keys(): - if name.startswith(backbone_prefixes): + if name.startswith(BACKBONE_PREFIXES) and not name.startswith(LLM_FROZEN_PREFIX): keys_to_load.append(name) elif load_delivery_head and name.startswith("head.delivery_head."): keys_to_load.append(name) @@ -410,45 +373,49 @@ class LNPModel(nn.Module): k: v for k, v in pretrain_state_dict.items() if k in keys_to_load } - missing, unexpected = [], [] + unexpected = [] model_state = self.state_dict() for k, v in filtered_state_dict.items(): - if k in model_state: - if model_state[k].shape == v.shape: - model_state[k] = v - else: - unexpected.append( - f"{k} (shape mismatch: {model_state[k].shape} vs {v.shape})" - ) + if k in model_state and model_state[k].shape == v.shape: + model_state[k] = v else: unexpected.append(k) + self.load_state_dict(model_state, strict=False) - if strict and (missing or unexpected): - raise RuntimeError(f"Missing keys: {missing}, Unexpected keys: {unexpected}") + if strict and unexpected: + raise RuntimeError(f"Unexpected keys: {unexpected}") class LNPModelWithoutMPNN(LNPModel): - """不使用 MPNN 的简化版本""" + """不使用 MPNN 的简化版本(化学 token 为 3 个)""" def __init__( self, d_model: int = 256, num_heads: int = 8, n_attn_layers: int = 4, + set_transformer_block: str = "sab", fusion_strategy: PoolingStrategy = "attention", head_hidden_dim: int = 128, dropout: float = 0.1, input_dims: Optional[Dict[str, int]] = None, - # ============ MoE 相关(新增) ============ + # ============ MoE 相关 ============ use_moe: bool = False, moe_n_experts: int = 4, moe_top_k: int = 2, moe_expert_hidden_mult: int = 2, moe_jitter_noise: float = 0.0, + # ============ LLM 相关 ============ + use_llm: bool = False, + llm_model_path: str = DEFAULT_MOLT5_PATH, + llm_freeze: bool = True, + llm_use_lora: bool = False, + llm_lora_r: int = 8, + llm_lora_alpha: int = 16, + llm_lora_dropout: float = 0.05, ) -> None: - # 移除 mpnn 维度 dims = input_dims or DEFAULT_INPUT_DIMS.copy() dims.pop("mpnn", None) @@ -456,6 +423,7 @@ class LNPModelWithoutMPNN(LNPModel): d_model=d_model, num_heads=num_heads, n_attn_layers=n_attn_layers, + set_transformer_block=set_transformer_block, fusion_strategy=fusion_strategy, head_hidden_dim=head_hidden_dim, dropout=dropout, @@ -467,5 +435,11 @@ class LNPModelWithoutMPNN(LNPModel): moe_top_k=moe_top_k, moe_expert_hidden_mult=moe_expert_hidden_mult, moe_jitter_noise=moe_jitter_noise, - ) - + use_llm=use_llm, + llm_model_path=llm_model_path, + llm_freeze=llm_freeze, + llm_use_lora=llm_use_lora, + llm_lora_r=llm_lora_r, + llm_lora_alpha=llm_lora_alpha, + llm_lora_dropout=llm_lora_dropout, + ) \ No newline at end of file diff --git a/lnp_ml/modeling/nested_cv_optuna.py b/lnp_ml/modeling/nested_cv_optuna.py index 0fc546a..fb805de 100644 --- a/lnp_ml/modeling/nested_cv_optuna.py +++ b/lnp_ml/modeling/nested_cv_optuna.py @@ -1,933 +1,1021 @@ -""" -嵌套交叉验证 + Optuna 超参调优 - -外层 5-fold StratifiedKFold(20% test / 80% train) -内层 3-fold StratifiedKFold(在 80% 上做 Optuna 超参搜索) - -使用方法: - python -m lnp_ml.modeling.nested_cv_optuna - -或通过 Makefile: - make nested_cv_tune DEVICE=cuda -""" - -import json -import math -from datetime import datetime -from pathlib import Path -from typing import Dict, List, Optional, Tuple, Union - -import numpy as np -import pandas as pd -import torch -from torch.utils.data import DataLoader, Subset -from sklearn.model_selection import StratifiedKFold -from loguru import logger -import typer - -try: - import optuna - from optuna.samplers import TPESampler -except ImportError: - optuna = None - TPESampler = None - -from lnp_ml.config import MODELS_DIR, INTERIM_DATA_DIR -from lnp_ml.dataset import ( - LNPDataset, - collate_fn, - process_dataframe, - TARGET_CLASSIFICATION_PDI, - TARGET_CLASSIFICATION_EE, - TARGET_TOXIC, -) -from tqdm import tqdm - -from lnp_ml.modeling.models import LNPModel, LNPModelWithoutMPNN -from lnp_ml.modeling.encoders.rdkit_encoder import CachedRDKitEncoder -from lnp_ml.modeling.trainer_balanced import ( - ClassWeights, - LossWeightsBalanced, - compute_class_weights_from_loader, - train_with_early_stopping, - train_fixed_epochs, - validate_balanced, -) - -# MPNN ensemble 默认路径 -DEFAULT_MPNN_ENSEMBLE_DIR = MODELS_DIR / "mpnn" / "all_amine_split_for_LiON" - -app = typer.Typer() - - -# ============ CompositeStrata 复合分层标签 ============ - -def build_composite_strata( - df: pd.DataFrame, - min_stratum_count: int = 5, -) -> Tuple[np.ndarray, Dict]: - """ - 构建复合分层标签(toxic × PDI × EE)。 - - Args: - df: 处理后的 DataFrame - min_stratum_count: 每个 stratum 最少样本数,低于此值合并为 RARE - - Returns: - (strata_array, strata_info) - - strata_array: 每个样本的 stratum 编码(整数) - - strata_info: 统计信息 - """ - n = len(df) - strata_labels = [] - - for i in range(n): - # Toxic stratum - if TARGET_TOXIC in df.columns: - toxic_val = df[TARGET_TOXIC].iloc[i] - if pd.notna(toxic_val) and toxic_val >= 0: - toxic_str = str(int(toxic_val)) - else: - toxic_str = "NA" - else: - toxic_str = "NA" - - # PDI stratum - if all(col in df.columns for col in TARGET_CLASSIFICATION_PDI): - pdi_vals = df[TARGET_CLASSIFICATION_PDI].iloc[i].values - if pdi_vals.sum() > 0: - pdi_str = str(int(np.argmax(pdi_vals))) - else: - pdi_str = "NA" - else: - pdi_str = "NA" - - # EE stratum - if all(col in df.columns for col in TARGET_CLASSIFICATION_EE): - ee_vals = df[TARGET_CLASSIFICATION_EE].iloc[i].values - if ee_vals.sum() > 0: - ee_str = str(int(np.argmax(ee_vals))) - else: - ee_str = "NA" - else: - ee_str = "NA" - - strata_labels.append(f"T{toxic_str}|P{pdi_str}|E{ee_str}") - - # 统计各 stratum 的样本数 - unique_strata, counts = np.unique(strata_labels, return_counts=True) - strata_counts = dict(zip(unique_strata, counts)) - - # 将稀疏 strata 合并为 RARE - rare_strata = [s for s, c in strata_counts.items() if c < min_stratum_count] - - final_labels = [] - for label in strata_labels: - if label in rare_strata: - final_labels.append("RARE") - else: - final_labels.append(label) - - # 编码为整数 - unique_final, encoded = np.unique(final_labels, return_inverse=True) - - strata_info = { - "original_strata_counts": strata_counts, - "rare_strata": rare_strata, - "final_strata": list(unique_final), - "final_strata_counts": dict(zip(*np.unique(final_labels, return_counts=True))), - "n_rare_merged": sum(strata_counts[s] for s in rare_strata) if rare_strata else 0, - } - - logger.info(f"Built composite strata: {len(unique_final)} unique strata") - logger.info(f" Rare strata merged: {len(rare_strata)} types, {strata_info['n_rare_merged']} samples") - - return encoded.astype(np.int64), strata_info - - -# ============ RDKit 缓存预热 ============ - -def warmup_rdkit_cache(smiles_list: List[str], batch_size: int = 256) -> Dict: - """预热 RDKit 特征缓存,返回可跨模型共享的缓存字典。""" - encoder = CachedRDKitEncoder() - unique_smiles = list(set(smiles_list)) - logger.info(f"Warming up RDKit cache for {len(unique_smiles)} unique SMILES...") - for i in tqdm(range(0, len(unique_smiles), batch_size), desc="Cache warmup"): - batch = unique_smiles[i:i + batch_size] - encoder(batch) - logger.success(f"Cache warmup complete. Cached {len(encoder._cache)} SMILES.") - return encoder._cache - - -# ============ 模型创建 ============ - -def find_mpnn_ensemble_paths(base_dir: Path = DEFAULT_MPNN_ENSEMBLE_DIR) -> List[str]: - """自动查找 MPNN ensemble 的 model.pt 文件。""" - model_paths = sorted(base_dir.glob("cv_*/fold_*/model_*/model.pt")) - if not model_paths: - raise FileNotFoundError(f"No model.pt files found in {base_dir}") - return [str(p) for p in model_paths] - - -def create_model( - d_model: int = 256, - num_heads: int = 8, - n_attn_layers: int = 4, - fusion_strategy: str = "attention", - head_hidden_dim: int = 128, - 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( - d_model=d_model, - num_heads=num_heads, - n_attn_layers=n_attn_layers, - fusion_strategy=fusion_strategy, - head_hidden_dim=head_hidden_dim, - dropout=dropout, - mpnn_ensemble_paths=ensemble_paths, - mpnn_device=mpnn_device, - **moe_kwargs, - ) - else: - return LNPModelWithoutMPNN( - d_model=d_model, - num_heads=num_heads, - n_attn_layers=n_attn_layers, - fusion_strategy=fusion_strategy, - head_hidden_dim=head_hidden_dim, - dropout=dropout, - **moe_kwargs, - ) - - -# ============ 评估指标 ============ - -def evaluate_on_test( - model: torch.nn.Module, - test_loader: DataLoader, - device: torch.device, -) -> Dict: - """在测试集上评估模型""" - from scipy.special import rel_entr - from sklearn.metrics import ( - mean_squared_error, - mean_absolute_error, - r2_score, - accuracy_score, - precision_score, - recall_score, - f1_score, - ) - - model.eval() - - preds = { - "size": [], "delivery": [], "pdi": [], "ee": [], "toxic": [], "biodist": [] - } - targets = { - "size": [], "delivery": [], "pdi": [], "ee": [], "toxic": [], "biodist": [] - } - - with torch.no_grad(): - for batch in test_loader: - smiles = batch["smiles"] - tabular = {k: v.to(device) for k, v in batch["tabular"].items()} - tgts = batch["targets"] - masks = batch["mask"] - - outputs = model(smiles, tabular) - - # 收集预测和真实值 - for task in ["size", "delivery"]: - if task in masks and masks[task].any(): - m = masks[task] - key = task if task == "size" else "delivery" - preds[task].extend(outputs[key].squeeze(-1)[m].cpu().numpy().tolist()) - targets[task].extend(tgts[key][m].cpu().numpy().tolist()) - - for task in ["pdi", "ee", "toxic"]: - if task in masks and masks[task].any(): - m = masks[task] - preds[task].extend(outputs[task][m].argmax(dim=-1).cpu().numpy().tolist()) - targets[task].extend(tgts[task][m].cpu().numpy().tolist()) - - if "biodist" in masks and masks["biodist"].any(): - m = masks["biodist"] - preds["biodist"].extend(outputs["biodist"][m].cpu().numpy().tolist()) - targets["biodist"].extend(tgts["biodist"][m].cpu().numpy().tolist()) - - # 计算指标 - results = {} - - # 回归任务 - for task in ["size", "delivery"]: - if preds[task]: - p = np.array(preds[task]) - t = np.array(targets[task]) - results[task] = { - "n_samples": len(p), - "mse": float(mean_squared_error(t, p)), - "rmse": float(np.sqrt(mean_squared_error(t, p))), - "mae": float(mean_absolute_error(t, p)), - "r2": float(r2_score(t, p)), - } - - # 分类任务 - for task in ["pdi", "ee", "toxic"]: - if preds[task]: - p = np.array(preds[task]) - t = np.array(targets[task]) - results[task] = { - "n_samples": len(p), - "accuracy": float(accuracy_score(t, p)), - "precision": float(precision_score(t, p, average="macro", zero_division=0)), - "recall": float(recall_score(t, p, average="macro", zero_division=0)), - "f1": float(f1_score(t, p, average="macro", zero_division=0)), - } - - # 分布任务 - if preds["biodist"]: - p = np.array(preds["biodist"]) - t = np.array(targets["biodist"]) - - def kl_divergence(p_arr, q_arr, eps=1e-10): - p_arr = np.clip(p_arr, eps, 1.0) - q_arr = np.clip(q_arr, eps, 1.0) - return float(np.sum(rel_entr(p_arr, q_arr), axis=-1).mean()) - - def js_divergence(p_arr, q_arr, eps=1e-10): - p_arr = np.clip(p_arr, eps, 1.0) - q_arr = np.clip(q_arr, eps, 1.0) - m = 0.5 * (p_arr + q_arr) - return float(0.5 * (np.sum(rel_entr(p_arr, m), axis=-1) + np.sum(rel_entr(q_arr, m), axis=-1)).mean()) - - results["biodist"] = { - "n_samples": len(p), - "kl_divergence": kl_divergence(t, p), - "js_divergence": js_divergence(t, p), - } - - return results - - -# ============ 预训练权重加载 ============ - -def load_pretrain_weights_to_model( - model: Union[LNPModel, LNPModelWithoutMPNN], - pretrain_state_dict: Dict, - d_model: int, - pretrain_config: Dict, - load_delivery_head: bool = True, -) -> bool: - """ - 加载预训练权重到模型。 - - Returns: - 是否成功加载 - """ - if pretrain_config.get("d_model") != d_model: - logger.warning( - f"d_model mismatch: pretrain={pretrain_config.get('d_model')}, " - f"current={d_model}. Skipping pretrain loading." - ) - return False - - model.load_pretrain_weights( - pretrain_state_dict=pretrain_state_dict, - load_delivery_head=load_delivery_head, - strict=False, - ) - return True - - -# ============ 内层 Optuna 调参 ============ -def run_inner_optuna( - full_dataset: LNPDataset, - inner_train_indices: np.ndarray, - strata: np.ndarray, - device: torch.device, - n_trials: int = 20, - epochs_per_trial: int = 30, - patience: int = 10, - batch_size: int = 32, - n_inner_folds: int = 3, - use_mpnn: bool = False, - seed: int = 42, - study_path: Optional[Path] = None, - pretrain_state_dict: Optional[Dict] = None, - 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 超参搜索。 - - Args: - full_dataset: 完整数据集 - inner_train_indices: 内层训练数据的索引(相对于 full_dataset) - strata: 每个样本的分层标签 - device: 设备 - n_trials: Optuna 试验数 - epochs_per_trial: 每个试验的最大 epoch - patience: 早停耐心值 - batch_size: 批次大小 - n_inner_folds: 内层折数 - use_mpnn: 是否使用 MPNN - seed: 随机种子 - study_path: 可选的 study 持久化路径 - pretrain_state_dict: 预训练权重 - pretrain_config: 预训练配置 - load_delivery_head: 是否加载 delivery head 权重 - - Returns: - (best_params, epoch_mean, study) - """ - if optuna is None: - raise ImportError("Optuna not installed. Run: pip install optuna") - - inner_strata = strata[inner_train_indices] - - # 固定架构参数(与预训练一致,确保权重完整加载) - _cfg = pretrain_config or {} - fixed_d_model = _cfg.get("d_model", 256) - fixed_num_heads = _cfg.get("num_heads", 8) - fixed_n_attn_layers = _cfg.get("n_attn_layers", 4) - fixed_fusion_strategy = _cfg.get("fusion_strategy", "attention") - fixed_head_hidden_dim = _cfg.get("head_hidden_dim", 128) - logger.info( - f"Fixed architecture params: d_model={fixed_d_model}, num_heads={fixed_num_heads}, " - f"n_attn_layers={fixed_n_attn_layers}, fusion={fixed_fusion_strategy}, " - f"head_hidden_dim={fixed_head_hidden_dim}" - ) - - def objective(trial: optuna.Trial) -> float: - d_model = fixed_d_model - num_heads = fixed_num_heads - n_attn_layers = fixed_n_attn_layers - fusion_strategy = fixed_fusion_strategy - head_hidden_dim = fixed_head_hidden_dim - - # 搜索训练超参数 - dropout = trial.suggest_float("dropout", 0.1, 0.5) - lr = trial.suggest_float("lr", 1e-5, 1e-3, log=True) - weight_decay = trial.suggest_float("weight_decay", 1e-5, 1e-1, log=True) - backbone_lr_ratio = trial.suggest_float("backbone_lr_ratio", 0.01, 1.0, log=True) - - # 内层 3-fold CV - inner_cv = StratifiedKFold( - n_splits=n_inner_folds, shuffle=True, random_state=seed - ) - - fold_val_losses = [] - fold_best_epochs = [] - - for inner_fold, (inner_train_idx, inner_val_idx) in enumerate( - inner_cv.split(inner_train_indices, inner_strata) - ): - # 获取实际的数据集索引 - actual_train_idx = inner_train_indices[inner_train_idx] - actual_val_idx = inner_train_indices[inner_val_idx] - - # 创建 DataLoader - train_subset = Subset(full_dataset, actual_train_idx.tolist()) - val_subset = Subset(full_dataset, actual_val_idx.tolist()) - - train_loader = DataLoader( - train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn - ) - val_loader = DataLoader( - val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn - ) - - # 计算类权重 - class_weights = compute_class_weights_from_loader(train_loader) - - # 创建模型 - model = create_model( - d_model=d_model, - num_heads=num_heads, - n_attn_layers=n_attn_layers, - fusion_strategy=fusion_strategy, - head_hidden_dim=head_hidden_dim, - 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 - - # 加载预训练权重 - if pretrain_state_dict is not None and pretrain_config is not None: - load_pretrain_weights_to_model( - model, pretrain_state_dict, d_model, pretrain_config, load_delivery_head - ) - - # 训练(带早停) - result = train_with_early_stopping( - model=model, - train_loader=train_loader, - val_loader=val_loader, - device=device, - lr=lr, - weight_decay=weight_decay, - epochs=epochs_per_trial, - patience=patience, - class_weights=class_weights, - backbone_lr_ratio=backbone_lr_ratio, - ) - - fold_val_losses.append(result["best_val_loss"]) - fold_best_epochs.append(result["best_epoch"]) - - # 记录 epoch_mean 到 trial - epoch_mean = int(round(np.mean(fold_best_epochs))) - trial.set_user_attr("epoch_mean", epoch_mean) - trial.set_user_attr("fold_best_epochs", fold_best_epochs) - - return np.mean(fold_val_losses) - - # 创建 study - storage = None - if study_path is not None: - storage = f"sqlite:///{study_path}" - - study = optuna.create_study( - direction="minimize", - sampler=TPESampler(seed=seed), - storage=storage, - study_name="inner_optuna", - load_if_exists=True, - ) - - study.optimize(objective, n_trials=n_trials, show_progress_bar=True) - - best_params = dict(study.best_trial.params) - best_params.update({ - "d_model": fixed_d_model, - "num_heads": fixed_num_heads, - "n_attn_layers": fixed_n_attn_layers, - "fusion_strategy": fixed_fusion_strategy, - "head_hidden_dim": fixed_head_hidden_dim, - }) - epoch_mean = study.best_trial.user_attrs.get("epoch_mean", epochs_per_trial) - - logger.info(f"Best trial: {study.best_trial.number}") - logger.info(f"Best val_loss: {study.best_trial.value:.4f}") - logger.info(f"Best params: {best_params}") - logger.info(f"Epoch mean: {epoch_mean}") - - return best_params, epoch_mean, study - - -# ============ 单 fold 执行(可跨进程调用) ============ -def _run_single_outer_fold( - outer_fold: int, - outer_train_idx: np.ndarray, - outer_test_idx: np.ndarray, - df: pd.DataFrame, - strata: np.ndarray, - fold_dir: Path, - n_trials: int, - epochs_per_trial: int, - inner_patience: int, - batch_size: int, - n_inner_folds: int, - use_mpnn: bool, - seed: int, - pretrain_state_dict: Optional[Dict], - 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 的完整流程(内层调参 + 外层训练 + 评估)。 - - 所有参数均为可序列化类型,以支持 spawn 多进程。 - """ - device = torch.device(device_str) - fold_dir = Path(fold_dir) - fold_dir.mkdir(parents=True, exist_ok=True) - - full_dataset = LNPDataset(df) - - logger.info(f"\n{'='*60}") - logger.info(f"OUTER FOLD {outer_fold}") - logger.info(f"{'='*60}") - logger.info(f"Train: {len(outer_train_idx)}, Test: {len(outer_test_idx)}") - - # 预热 RDKit 缓存(在整个 fold 内共享) - rdkit_cache = warmup_rdkit_cache(full_dataset.smiles) - - # 保存 split indices - with open(fold_dir / "splits.json", "w") as f: - json.dump({ - "outer_train_idx": outer_train_idx.tolist(), - "outer_test_idx": outer_test_idx.tolist(), - }, f) - - # 内层 Optuna 调参 - logger.info(f"\nRunning inner Optuna with {n_trials} trials...") - study_path = fold_dir / "optuna_study.sqlite3" - - best_params, epoch_mean, study = run_inner_optuna( - full_dataset=full_dataset, - inner_train_indices=outer_train_idx, - strata=strata, - device=device, - n_trials=n_trials, - epochs_per_trial=epochs_per_trial, - patience=inner_patience, - batch_size=batch_size, - n_inner_folds=n_inner_folds, - use_mpnn=use_mpnn, - seed=seed + outer_fold, - study_path=study_path, - pretrain_state_dict=pretrain_state_dict, - 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, - ) - - # 保存最佳参数 - with open(fold_dir / "best_params.json", "w") as f: - json.dump(best_params, f, indent=2) - - with open(fold_dir / "epoch_mean.json", "w") as f: - json.dump({"epoch_mean": epoch_mean}, f) - - # 外层训练(使用最优超参,固定 epoch 数,不 early-stop) - logger.info(f"\nTraining outer fold with best params, epochs={epoch_mean}...") - - train_subset = Subset(full_dataset, outer_train_idx.tolist()) - test_subset = Subset(full_dataset, outer_test_idx.tolist()) - - train_loader = DataLoader( - train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn - ) - test_loader = DataLoader( - test_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn - ) - - class_weights = compute_class_weights_from_loader(train_loader) - - model = create_model( - d_model=best_params["d_model"], - num_heads=best_params["num_heads"], - n_attn_layers=best_params["n_attn_layers"], - fusion_strategy=best_params["fusion_strategy"], - head_hidden_dim=best_params["head_hidden_dim"], - 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 - - if pretrain_state_dict is not None and pretrain_config is not None: - loaded = load_pretrain_weights_to_model( - model, pretrain_state_dict, best_params["d_model"], - pretrain_config, load_delivery_head - ) - if loaded: - logger.info(f"Loaded pretrain weights for outer fold {outer_fold}") - - train_result = train_fixed_epochs( - model=model, - train_loader=train_loader, - val_loader=None, - device=device, - lr=best_params["lr"], - weight_decay=best_params["weight_decay"], - epochs=epoch_mean, - class_weights=class_weights, - use_cosine_annealing=True, - backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0), - ) - - model.load_state_dict(train_result["final_state"]) - model = model.to(device) - - config = { - "d_model": best_params["d_model"], - "num_heads": best_params["num_heads"], - "n_attn_layers": best_params["n_attn_layers"], - "fusion_strategy": best_params["fusion_strategy"], - "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, - "epoch_mean": epoch_mean, - "best_params": best_params, - }, fold_dir / "model.pt") - - with open(fold_dir / "history.json", "w") as f: - json.dump(train_result["history"], f, indent=2) - - # 在测试集上评估 - logger.info("Evaluating on outer test set...") - test_metrics = evaluate_on_test(model, test_loader, device) - - with open(fold_dir / "test_metrics.json", "w") as f: - json.dump(test_metrics, f, indent=2) - - logger.info(f"\nOuter Fold {outer_fold} Test Results:") - for task, metrics in test_metrics.items(): - if "rmse" in metrics: - logger.info(f" {task}: RMSE={metrics['rmse']:.4f}, R²={metrics['r2']:.4f}") - elif "accuracy" in metrics: - logger.info(f" {task}: Acc={metrics['accuracy']:.4f}, F1={metrics['f1']:.4f}") - elif "kl_divergence" in metrics: - logger.info(f" {task}: KL={metrics['kl_divergence']:.4f}, JS={metrics['js_divergence']:.4f}") - - return { - "fold": outer_fold, - "best_params": best_params, - "epoch_mean": epoch_mean, - "test_metrics": test_metrics, - } - - -# ============ 主流程 ============ -@app.command() -def main( - input_path: Path = INTERIM_DATA_DIR / "internal.csv", - output_dir: Path = MODELS_DIR / "nested_cv", - # CV 参数 - n_outer_folds: int = 5, - n_inner_folds: int = 3, - min_stratum_count: int = 5, - seed: int = 42, - # Optuna 参数 - n_trials: int = 20, - epochs_per_trial: int = 30, - inner_patience: int = 10, - # 训练参数 - batch_size: int = 32, - # 预训练权重 - init_from_pretrain: Optional[Path] = None, - 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, - # 设备 - device: str = "cuda" if torch.cuda.is_available() else "cpu", -): - """ - 嵌套交叉验证 + Optuna 超参调优。 - - 外层 5-fold(20% test / 80% train),内层 3-fold Optuna 调参。 - 外层训练不使用 early-stopping,epoch 数使用内层 best trial 的 epoch_mean。 - - 使用 --init-from-pretrain 从预训练 checkpoint 初始化模型权重。 - 使用 --parallel 同时运行所有外层 fold(需要足够 GPU 显存)。 - """ - if optuna is None: - logger.error("Optuna not installed. Run: pip install optuna") - raise typer.Exit(1) - - logger.info(f"Using device: {device}") - device = torch.device(device) - - # 加载预训练权重(如果指定) - pretrain_state_dict = None - pretrain_config = None - if init_from_pretrain is not None: - if init_from_pretrain.exists(): - logger.info(f"Loading pretrain weights from {init_from_pretrain}") - checkpoint = torch.load(init_from_pretrain, map_location="cpu", weights_only=False) - pretrain_state_dict = checkpoint["model_state_dict"] - pretrain_config = checkpoint.get("config", {}) - logger.success(f"Loaded pretrain checkpoint (d_model={pretrain_config.get('d_model')})") - else: - logger.warning(f"Pretrain checkpoint not found: {init_from_pretrain}, skipping") - - # 创建输出目录(带时间戳) - run_name = datetime.now().strftime("%Y%m%d_%H%M%S") - run_dir = output_dir / run_name - run_dir.mkdir(parents=True, exist_ok=True) - logger.info(f"Output directory: {run_dir}") - - # 加载数据 - logger.info(f"Loading data from {input_path}") - df = pd.read_csv(input_path) - logger.info(f"Loaded {len(df)} samples") - - # 处理数据 - logger.info("Processing dataframe...") - df = process_dataframe(df) - - # 构建复合分层标签 - logger.info("Building composite strata...") - strata, strata_info = build_composite_strata(df, min_stratum_count) - - # 保存 strata 信息 - with open(run_dir / "strata_info.json", "w") as f: - json.dump(strata_info, f, indent=2, default=str) - - # 创建完整数据集(仅用于获取样本数做 split) - n_samples = len(LNPDataset(df)) - - # 外层 CV split - outer_cv = StratifiedKFold( - n_splits=n_outer_folds, shuffle=True, random_state=seed - ) - - device_str = str(device) - fold_args = [] - for outer_fold, (outer_train_idx, outer_test_idx) in enumerate( - outer_cv.split(np.arange(n_samples), strata) - ): - fold_args.append(dict( - outer_fold=outer_fold, - outer_train_idx=outer_train_idx, - outer_test_idx=outer_test_idx, - df=df, - strata=strata, - fold_dir=run_dir / f"outer_fold_{outer_fold}", - n_trials=n_trials, - epochs_per_trial=epochs_per_trial, - inner_patience=inner_patience, - batch_size=batch_size, - n_inner_folds=n_inner_folds, - use_mpnn=use_mpnn, - seed=seed, - pretrain_state_dict=pretrain_state_dict, - 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: - import multiprocessing as mp - from concurrent.futures import ProcessPoolExecutor - - ctx = mp.get_context("spawn") - logger.info(f"Running {n_outer_folds} outer folds in PARALLEL (spawn)") - with ProcessPoolExecutor(max_workers=n_outer_folds, mp_context=ctx) as executor: - futures = [executor.submit(_run_single_outer_fold, **args) for args in fold_args] - outer_results = [f.result() for f in futures] - outer_results.sort(key=lambda r: r["fold"]) - else: - logger.info(f"Running {n_outer_folds} outer folds SEQUENTIALLY") - outer_results = [] - for args in fold_args: - result = _run_single_outer_fold(**args) - outer_results.append(result) - - # 汇总结果 - logger.info("\n" + "=" * 60) - logger.info("NESTED CV COMPLETE") - logger.info("=" * 60) - - # 计算汇总统计 - summary = {"fold_results": outer_results} - - # 对每个任务计算均值和标准差 - tasks_with_metrics = {} - for result in outer_results: - for task, metrics in result["test_metrics"].items(): - if task not in tasks_with_metrics: - tasks_with_metrics[task] = {k: [] for k in metrics.keys() if k != "n_samples"} - for k, v in metrics.items(): - if k != "n_samples": - tasks_with_metrics[task][k].append(v) - - summary["summary_stats"] = {} - for task, metrics_dict in tasks_with_metrics.items(): - summary["summary_stats"][task] = {} - for metric_name, values in metrics_dict.items(): - summary["summary_stats"][task][f"{metric_name}_mean"] = float(np.mean(values)) - summary["summary_stats"][task][f"{metric_name}_std"] = float(np.std(values)) - - # 打印汇总 - logger.info("\n[Summary Statistics]") - for task, stats in summary["summary_stats"].items(): - if "rmse_mean" in stats: - logger.info( - f" {task}: RMSE={stats['rmse_mean']:.4f}±{stats['rmse_std']:.4f}, " - f"R²={stats['r2_mean']:.4f}±{stats['r2_std']:.4f}" - ) - elif "accuracy_mean" in stats: - logger.info( - f" {task}: Acc={stats['accuracy_mean']:.4f}±{stats['accuracy_std']:.4f}, " - f"F1={stats['f1_mean']:.4f}±{stats['f1_std']:.4f}" - ) - elif "kl_divergence_mean" in stats: - logger.info( - f" {task}: KL={stats['kl_divergence_mean']:.4f}±{stats['kl_divergence_std']:.4f}, " - f"JS={stats['js_divergence_mean']:.4f}±{stats['js_divergence_std']:.4f}" - ) - - # 保存汇总 - with open(run_dir / "summary.json", "w") as f: - json.dump(summary, f, indent=2) - - logger.success(f"\nAll results saved to {run_dir}") - - -if __name__ == "__main__": - app() - +""" +嵌套交叉验证 + Optuna 超参调优 + +外层 5-fold StratifiedKFold(20% test / 80% train) +内层 3-fold StratifiedKFold(在 80% 上做 Optuna 超参搜索) + +使用方法: + python -m lnp_ml.modeling.nested_cv_optuna + +或通过 Makefile: + make nested_cv_tune DEVICE=cuda +""" + +import json +import math +from datetime import datetime +from pathlib import Path +from typing import Dict, List, Optional, Tuple, Union + +import numpy as np +import pandas as pd +import torch +from torch.utils.data import DataLoader, Subset +from sklearn.model_selection import StratifiedKFold +from loguru import logger +import typer + +try: + import optuna + from optuna.samplers import TPESampler +except ImportError: + optuna = None + TPESampler = None + +from lnp_ml.config import MODELS_DIR, INTERIM_DATA_DIR +from lnp_ml.dataset import ( + LNPDataset, + collate_fn, + process_dataframe, + TARGET_CLASSIFICATION_PDI, + TARGET_CLASSIFICATION_EE, + TARGET_TOXIC, +) +from tqdm import tqdm + +from lnp_ml.modeling.models import LNPModel, LNPModelWithoutMPNN +from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH +from lnp_ml.utils.seed import set_global_seed +from lnp_ml.modeling.encoders.rdkit_encoder import CachedRDKitEncoder +from lnp_ml.modeling.trainer_balanced import ( + ClassWeights, + LossWeightsBalanced, + compute_class_weights_from_loader, + train_with_early_stopping, + train_fixed_epochs, + validate_balanced, +) + +# MPNN ensemble 默认路径 +DEFAULT_MPNN_ENSEMBLE_DIR = MODELS_DIR / "mpnn" / "all_amine_split_for_LiON" + +app = typer.Typer() + + +# ============ CompositeStrata 复合分层标签 ============ + +def build_composite_strata( + df: pd.DataFrame, + min_stratum_count: int = 5, +) -> Tuple[np.ndarray, Dict]: + """ + 构建复合分层标签(toxic × PDI × EE)。 + + Args: + df: 处理后的 DataFrame + min_stratum_count: 每个 stratum 最少样本数,低于此值合并为 RARE + + Returns: + (strata_array, strata_info) + - strata_array: 每个样本的 stratum 编码(整数) + - strata_info: 统计信息 + """ + n = len(df) + strata_labels = [] + + for i in range(n): + # Toxic stratum + if TARGET_TOXIC in df.columns: + toxic_val = df[TARGET_TOXIC].iloc[i] + if pd.notna(toxic_val) and toxic_val >= 0: + toxic_str = str(int(toxic_val)) + else: + toxic_str = "NA" + else: + toxic_str = "NA" + + # PDI stratum + if all(col in df.columns for col in TARGET_CLASSIFICATION_PDI): + pdi_vals = df[TARGET_CLASSIFICATION_PDI].iloc[i].values + if pdi_vals.sum() > 0: + pdi_str = str(int(np.argmax(pdi_vals))) + else: + pdi_str = "NA" + else: + pdi_str = "NA" + + # EE stratum + if all(col in df.columns for col in TARGET_CLASSIFICATION_EE): + ee_vals = df[TARGET_CLASSIFICATION_EE].iloc[i].values + if ee_vals.sum() > 0: + ee_str = str(int(np.argmax(ee_vals))) + else: + ee_str = "NA" + else: + ee_str = "NA" + + strata_labels.append(f"T{toxic_str}|P{pdi_str}|E{ee_str}") + + # 统计各 stratum 的样本数 + unique_strata, counts = np.unique(strata_labels, return_counts=True) + strata_counts = dict(zip(unique_strata, counts)) + + # 将稀疏 strata 合并为 RARE + rare_strata = [s for s, c in strata_counts.items() if c < min_stratum_count] + + final_labels = [] + for label in strata_labels: + if label in rare_strata: + final_labels.append("RARE") + else: + final_labels.append(label) + + # 编码为整数 + unique_final, encoded = np.unique(final_labels, return_inverse=True) + + strata_info = { + "original_strata_counts": strata_counts, + "rare_strata": rare_strata, + "final_strata": list(unique_final), + "final_strata_counts": dict(zip(*np.unique(final_labels, return_counts=True))), + "n_rare_merged": sum(strata_counts[s] for s in rare_strata) if rare_strata else 0, + } + + logger.info(f"Built composite strata: {len(unique_final)} unique strata") + logger.info(f" Rare strata merged: {len(rare_strata)} types, {strata_info['n_rare_merged']} samples") + + return encoded.astype(np.int64), strata_info + + +# ============ RDKit 缓存预热 ============ + +def warmup_rdkit_cache(smiles_list: List[str], batch_size: int = 256) -> Dict: + """预热 RDKit 特征缓存,返回可跨模型共享的缓存字典。""" + encoder = CachedRDKitEncoder() + unique_smiles = list(set(smiles_list)) + logger.info(f"Warming up RDKit cache for {len(unique_smiles)} unique SMILES...") + for i in tqdm(range(0, len(unique_smiles), batch_size), desc="Cache warmup"): + batch = unique_smiles[i:i + batch_size] + encoder(batch) + logger.success(f"Cache warmup complete. Cached {len(encoder._cache)} SMILES.") + return encoder._cache + + +# ============ 模型创建 ============ + +def find_mpnn_ensemble_paths(base_dir: Path = DEFAULT_MPNN_ENSEMBLE_DIR) -> List[str]: + """自动查找 MPNN ensemble 的 model.pt 文件。""" + model_paths = sorted(base_dir.glob("cv_*/fold_*/model_*/model.pt")) + if not model_paths: + raise FileNotFoundError(f"No model.pt files found in {base_dir}") + return [str(p) for p in model_paths] + + +def create_model( + d_model: int = 256, + num_heads: int = 8, + n_attn_layers: int = 4, + fusion_strategy: str = "attention", + head_hidden_dim: int = 128, + dropout: float = 0.1, + use_mpnn: bool = False, + mpnn_device: str = "cpu", + set_transformer_block: str = "sab", + # MoE + use_moe: bool = False, + moe_n_experts: int = 4, + moe_top_k: int = 2, + moe_expert_hidden_mult: int = 2, + moe_jitter_noise: float = 0.0, + # LLM(打包成 dict 以减少多进程参数传递) + llm_kwargs: Optional[Dict] = None, +) -> Union[LNPModel, LNPModelWithoutMPNN]: + """创建模型。llm_kwargs 含 use_llm / llm_model_path / llm_freeze / llm_use_lora / llm_lora_*。""" + extra_kwargs = dict( + set_transformer_block=set_transformer_block, + use_moe=use_moe, + moe_n_experts=moe_n_experts, + moe_top_k=moe_top_k, + moe_expert_hidden_mult=moe_expert_hidden_mult, + moe_jitter_noise=moe_jitter_noise, + **(llm_kwargs or {}), + ) + + if use_mpnn: + ensemble_paths = find_mpnn_ensemble_paths() + return LNPModel( + d_model=d_model, + num_heads=num_heads, + n_attn_layers=n_attn_layers, + fusion_strategy=fusion_strategy, + head_hidden_dim=head_hidden_dim, + dropout=dropout, + mpnn_ensemble_paths=ensemble_paths, + mpnn_device=mpnn_device, + **extra_kwargs, + ) + else: + return LNPModelWithoutMPNN( + d_model=d_model, + num_heads=num_heads, + n_attn_layers=n_attn_layers, + fusion_strategy=fusion_strategy, + head_hidden_dim=head_hidden_dim, + dropout=dropout, + **extra_kwargs, + ) + + +# ============ 评估指标 ============ + +def evaluate_on_test( + model: torch.nn.Module, + test_loader: DataLoader, + device: torch.device, +) -> Dict: + """在测试集上评估模型""" + from scipy.special import rel_entr + from sklearn.metrics import ( + mean_squared_error, + mean_absolute_error, + r2_score, + accuracy_score, + precision_score, + recall_score, + f1_score, + ) + + model.eval() + + preds = { + "size": [], "delivery": [], "pdi": [], "ee": [], "toxic": [], "biodist": [] + } + targets = { + "size": [], "delivery": [], "pdi": [], "ee": [], "toxic": [], "biodist": [] + } + + with torch.no_grad(): + for batch in test_loader: + smiles = batch["smiles"] + tabular = {k: v.to(device) for k, v in batch["tabular"].items()} + tgts = batch["targets"] + masks = batch["mask"] + + outputs = model(smiles, tabular) + + # 收集预测和真实值 + for task in ["size", "delivery"]: + if task in masks and masks[task].any(): + m = masks[task] + key = task if task == "size" else "delivery" + preds[task].extend(outputs[key].squeeze(-1)[m].cpu().numpy().tolist()) + targets[task].extend(tgts[key][m].cpu().numpy().tolist()) + + for task in ["pdi", "ee", "toxic"]: + if task in masks and masks[task].any(): + m = masks[task] + preds[task].extend(outputs[task][m].argmax(dim=-1).cpu().numpy().tolist()) + targets[task].extend(tgts[task][m].cpu().numpy().tolist()) + + if "biodist" in masks and masks["biodist"].any(): + m = masks["biodist"] + preds["biodist"].extend(outputs["biodist"][m].cpu().numpy().tolist()) + targets["biodist"].extend(tgts["biodist"][m].cpu().numpy().tolist()) + + # 计算指标 + results = {} + + # 回归任务 + for task in ["size", "delivery"]: + if preds[task]: + p = np.array(preds[task]) + t = np.array(targets[task]) + results[task] = { + "n_samples": len(p), + "mse": float(mean_squared_error(t, p)), + "rmse": float(np.sqrt(mean_squared_error(t, p))), + "mae": float(mean_absolute_error(t, p)), + "r2": float(r2_score(t, p)), + } + + # 分类任务 + for task in ["pdi", "ee", "toxic"]: + if preds[task]: + p = np.array(preds[task]) + t = np.array(targets[task]) + results[task] = { + "n_samples": len(p), + "accuracy": float(accuracy_score(t, p)), + "precision": float(precision_score(t, p, average="macro", zero_division=0)), + "recall": float(recall_score(t, p, average="macro", zero_division=0)), + "f1": float(f1_score(t, p, average="macro", zero_division=0)), + } + + # 分布任务 + if preds["biodist"]: + p = np.array(preds["biodist"]) + t = np.array(targets["biodist"]) + + def kl_divergence(p_arr, q_arr, eps=1e-10): + p_arr = np.clip(p_arr, eps, 1.0) + q_arr = np.clip(q_arr, eps, 1.0) + return float(np.sum(rel_entr(p_arr, q_arr), axis=-1).mean()) + + def js_divergence(p_arr, q_arr, eps=1e-10): + p_arr = np.clip(p_arr, eps, 1.0) + q_arr = np.clip(q_arr, eps, 1.0) + m = 0.5 * (p_arr + q_arr) + return float(0.5 * (np.sum(rel_entr(p_arr, m), axis=-1) + np.sum(rel_entr(q_arr, m), axis=-1)).mean()) + + results["biodist"] = { + "n_samples": len(p), + "kl_divergence": kl_divergence(t, p), + "js_divergence": js_divergence(t, p), + } + + return results + + +# ============ 预训练权重加载 ============ + +def load_pretrain_weights_to_model( + model: Union[LNPModel, LNPModelWithoutMPNN], + pretrain_state_dict: Dict, + d_model: int, + pretrain_config: Dict, + load_delivery_head: bool = True, +) -> bool: + """ + 加载预训练权重到模型。 + + Returns: + 是否成功加载 + """ + if pretrain_config.get("d_model") != d_model: + logger.warning( + f"d_model mismatch: pretrain={pretrain_config.get('d_model')}, " + f"current={d_model}. Skipping pretrain loading." + ) + return False + + model.load_pretrain_weights( + pretrain_state_dict=pretrain_state_dict, + load_delivery_head=load_delivery_head, + strict=False, + ) + return True + + +# ============ 内层 Optuna 调参 ============ +def run_inner_optuna( + full_dataset: LNPDataset, + inner_train_indices: np.ndarray, + strata: np.ndarray, + device: torch.device, + n_trials: int = 20, + epochs_per_trial: int = 30, + patience: int = 10, + batch_size: int = 32, + n_inner_folds: int = 3, + use_mpnn: bool = False, + seed: int = 42, + study_path: Optional[Path] = None, + pretrain_state_dict: Optional[Dict] = None, + 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, + set_transformer_block: str = "sab", + llm_kwargs: Optional[Dict] = None, +) -> Tuple[Dict, int, optuna.Study]: + """ + 在内层数据上运行 Optuna 超参搜索。 + + Args: + full_dataset: 完整数据集 + inner_train_indices: 内层训练数据的索引(相对于 full_dataset) + strata: 每个样本的分层标签 + device: 设备 + n_trials: Optuna 试验数 + epochs_per_trial: 每个试验的最大 epoch + patience: 早停耐心值 + batch_size: 批次大小 + n_inner_folds: 内层折数 + use_mpnn: 是否使用 MPNN + seed: 随机种子 + study_path: 可选的 study 持久化路径 + pretrain_state_dict: 预训练权重 + pretrain_config: 预训练配置 + load_delivery_head: 是否加载 delivery head 权重 + + Returns: + (best_params, epoch_mean, study) + """ + if optuna is None: + raise ImportError("Optuna not installed. Run: pip install optuna") + + inner_strata = strata[inner_train_indices] + + # 固定架构参数(与预训练一致,确保权重完整加载) + _cfg = pretrain_config or {} + fixed_d_model = _cfg.get("d_model", 256) + fixed_num_heads = _cfg.get("num_heads", 8) + fixed_n_attn_layers = _cfg.get("n_attn_layers", 4) + fixed_fusion_strategy = _cfg.get("fusion_strategy", "attention") + fixed_head_hidden_dim = _cfg.get("head_hidden_dim", 128) + logger.info( + f"Fixed architecture params: d_model={fixed_d_model}, num_heads={fixed_num_heads}, " + f"n_attn_layers={fixed_n_attn_layers}, fusion={fixed_fusion_strategy}, " + f"head_hidden_dim={fixed_head_hidden_dim}" + ) + + def objective(trial: optuna.Trial) -> float: + d_model = fixed_d_model + num_heads = fixed_num_heads + n_attn_layers = fixed_n_attn_layers + fusion_strategy = fixed_fusion_strategy + head_hidden_dim = fixed_head_hidden_dim + + # 搜索训练超参数 + dropout = trial.suggest_float("dropout", 0.1, 0.5) + lr = trial.suggest_float("lr", 1e-5, 1e-3, log=True) + weight_decay = trial.suggest_float("weight_decay", 1e-5, 1e-1, log=True) + backbone_lr_ratio = trial.suggest_float("backbone_lr_ratio", 0.01, 1.0, log=True) + + # MoE / LLM + use_llm = bool(llm_kwargs.get("use_llm", False)) + moe_ne_t = trial.suggest_categorical("moe_n_experts", [2, 4]) if use_moe else moe_n_experts + moe_tk_t = trial.suggest_int("moe_top_k", 1, 2) if use_moe else moe_top_k + moe_hm_t = trial.suggest_categorical("moe_expert_hidden_mult", [1, 2]) if use_moe else moe_expert_hidden_mult + llm_kwargs_t = dict(llm_kwargs) + if use_llm: + llm_kwargs_t["llm_use_lora"] = True + llm_kwargs_t["llm_lora_r"] = trial.suggest_categorical("llm_lora_r", [8, 16, 32]) + + # 内层 3-fold CV + inner_cv = StratifiedKFold( + n_splits=n_inner_folds, shuffle=True, random_state=seed + ) + + fold_val_losses = [] + fold_best_epochs = [] + + for inner_fold, (inner_train_idx, inner_val_idx) in enumerate( + inner_cv.split(inner_train_indices, inner_strata) + ): + # 获取实际的数据集索引 + actual_train_idx = inner_train_indices[inner_train_idx] + actual_val_idx = inner_train_indices[inner_val_idx] + + # 创建 DataLoader + train_subset = Subset(full_dataset, actual_train_idx.tolist()) + val_subset = Subset(full_dataset, actual_val_idx.tolist()) + + train_loader = DataLoader( + train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn + ) + val_loader = DataLoader( + val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn + ) + + # 计算类权重 + class_weights = compute_class_weights_from_loader(train_loader) + + # 创建模型 + model = create_model( + d_model=d_model, + num_heads=num_heads, + n_attn_layers=n_attn_layers, + fusion_strategy=fusion_strategy, + head_hidden_dim=head_hidden_dim, + dropout=dropout, + use_mpnn=use_mpnn, + mpnn_device=device.type, + use_moe=use_moe, + moe_n_experts=moe_ne_t, + moe_top_k=moe_tk_t, + moe_expert_hidden_mult=moe_hm_t, + moe_jitter_noise=moe_jitter_noise, + set_transformer_block=set_transformer_block, + llm_kwargs=llm_kwargs_t, + ) + if rdkit_cache is not None: + model.rdkit_encoder._cache = rdkit_cache + + # 加载预训练权重 + if pretrain_state_dict is not None and pretrain_config is not None: + load_pretrain_weights_to_model( + model, pretrain_state_dict, d_model, pretrain_config, load_delivery_head + ) + + # 训练(带早停) + result = train_with_early_stopping( + model=model, + train_loader=train_loader, + val_loader=val_loader, + device=device, + lr=lr, + weight_decay=weight_decay, + epochs=epochs_per_trial, + patience=patience, + class_weights=class_weights, + backbone_lr_ratio=backbone_lr_ratio, + ) + + fold_val_losses.append(result["best_val_loss"]) + fold_best_epochs.append(result["best_epoch"]) + + # 记录 epoch_mean 到 trial + epoch_mean = int(round(np.mean(fold_best_epochs))) + trial.set_user_attr("epoch_mean", epoch_mean) + trial.set_user_attr("fold_best_epochs", fold_best_epochs) + + return np.mean(fold_val_losses) + + # 创建 study + storage = None + if study_path is not None: + storage = f"sqlite:///{study_path}" + + study = optuna.create_study( + direction="minimize", + sampler=TPESampler(seed=seed), + storage=storage, + study_name="inner_optuna", + load_if_exists=True, + ) + + study.optimize(objective, n_trials=n_trials, show_progress_bar=True) + + best_params = dict(study.best_trial.params) + best_params.update({ + "d_model": fixed_d_model, + "num_heads": fixed_num_heads, + "n_attn_layers": fixed_n_attn_layers, + "fusion_strategy": fixed_fusion_strategy, + "head_hidden_dim": fixed_head_hidden_dim, + "set_transformer_block": set_transformer_block, + }) + epoch_mean = study.best_trial.user_attrs.get("epoch_mean", epochs_per_trial) + + logger.info(f"Best trial: {study.best_trial.number}") + logger.info(f"Best val_loss: {study.best_trial.value:.4f}") + logger.info(f"Best params: {best_params}") + logger.info(f"Epoch mean: {epoch_mean}") + + return best_params, epoch_mean, study + + +# ============ 单 fold 执行(可跨进程调用) ============ +def _run_single_outer_fold( + outer_fold: int, + outer_train_idx: np.ndarray, + outer_test_idx: np.ndarray, + df: pd.DataFrame, + strata: np.ndarray, + fold_dir: Path, + n_trials: int, + epochs_per_trial: int, + inner_patience: int, + batch_size: int, + n_inner_folds: int, + use_mpnn: bool, + seed: int, + pretrain_state_dict: Optional[Dict], + 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, + set_transformer_block: str = "sab", + llm_kwargs: Optional[Dict] = None, + precomputed_best_params: Optional[Dict] = None, + precomputed_epoch_mean: Optional[int] = None, +) -> Dict: + """ + 执行单个外层 fold 的完整流程(内层调参 + 外层训练 + 评估)。 + + 所有参数均为可序列化类型,以支持 spawn 多进程。 + """ + device = torch.device(device_str) + set_global_seed(seed + outer_fold) + fold_dir = Path(fold_dir) + fold_dir.mkdir(parents=True, exist_ok=True) + + full_dataset = LNPDataset(df) + + logger.info(f"\n{'='*60}") + logger.info(f"OUTER FOLD {outer_fold}") + logger.info(f"{'='*60}") + logger.info(f"Train: {len(outer_train_idx)}, Test: {len(outer_test_idx)}") + + # 预热 RDKit 缓存(在整个 fold 内共享) + rdkit_cache = warmup_rdkit_cache(full_dataset.smiles) + + # 保存 split indices + with open(fold_dir / "splits.json", "w") as f: + json.dump({ + "outer_train_idx": outer_train_idx.tolist(), + "outer_test_idx": outer_test_idx.tolist(), + }, f) + + # 内层 Optuna 调参(方式 B:已有超参则跳过,直接复用) + if precomputed_best_params is not None and precomputed_epoch_mean is not None: + logger.info("Reusing precomputed best_params (skip inner Optuna).") + best_params = dict(precomputed_best_params) + epoch_mean = int(precomputed_epoch_mean) + else: + logger.info(f"\nRunning inner Optuna with {n_trials} trials...") + study_path = fold_dir / "optuna_study.sqlite3" + best_params, epoch_mean, study = run_inner_optuna( + full_dataset=full_dataset, + inner_train_indices=outer_train_idx, + strata=strata, + device=device, + n_trials=n_trials, + epochs_per_trial=epochs_per_trial, + patience=inner_patience, + batch_size=batch_size, + n_inner_folds=n_inner_folds, + use_mpnn=use_mpnn, + seed=seed + outer_fold, + study_path=study_path, + pretrain_state_dict=pretrain_state_dict, + 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, + set_transformer_block=set_transformer_block, + llm_kwargs=llm_kwargs, + ) + + # 保存最佳参数 + with open(fold_dir / "best_params.json", "w") as f: + json.dump(best_params, f, indent=2) + with open(fold_dir / "epoch_mean.json", "w") as f: + json.dump({"epoch_mean": epoch_mean}, f) + + # 外层训练(使用最优超参,固定 epoch 数,不 early-stop) + logger.info(f"\nTraining outer fold with best params, epochs={epoch_mean}...") + + # 从 outer_train 再切 10% 作为"监控用 val"(只画曲线、看 gap,不参与选模型/早停) + _rng = np.random.RandomState(seed + outer_fold) + _perm = _rng.permutation(len(outer_train_idx)) + _n_val = max(1, int(0.1 * len(outer_train_idx))) + fit_idx = outer_train_idx[_perm[_n_val:]] + monitor_idx = outer_train_idx[_perm[:_n_val]] + + train_subset = Subset(full_dataset, fit_idx.tolist()) + monitor_subset = Subset(full_dataset, monitor_idx.tolist()) + test_subset = Subset(full_dataset, outer_test_idx.tolist()) + + train_loader = DataLoader( + train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn + ) + monitor_loader = DataLoader( + monitor_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn + ) + test_loader = DataLoader( + test_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn + ) + + class_weights = compute_class_weights_from_loader(train_loader) + + model = create_model( + d_model=best_params["d_model"], + num_heads=best_params["num_heads"], + n_attn_layers=best_params["n_attn_layers"], + fusion_strategy=best_params["fusion_strategy"], + head_hidden_dim=best_params["head_hidden_dim"], + dropout=best_params["dropout"], + use_mpnn=use_mpnn, + mpnn_device=device.type, + use_moe=use_moe, + moe_n_experts=best_params.get("moe_n_experts", moe_n_experts), + moe_top_k=best_params.get("moe_top_k", moe_top_k), + moe_expert_hidden_mult=best_params.get("moe_expert_hidden_mult", moe_expert_hidden_mult), + moe_jitter_noise=moe_jitter_noise, + set_transformer_block=best_params.get("set_transformer_block", "sab"), + llm_kwargs={**llm_kwargs, **({"llm_use_lora": True, "llm_lora_r": best_params["llm_lora_r"]} + if (llm_kwargs.get("use_llm", False) and "llm_lora_r" in best_params) else {})}, + ) + model.rdkit_encoder._cache = rdkit_cache + + if pretrain_state_dict is not None and pretrain_config is not None: + loaded = load_pretrain_weights_to_model( + model, pretrain_state_dict, best_params["d_model"], + pretrain_config, load_delivery_head + ) + if loaded: + logger.info(f"Loaded pretrain weights for outer fold {outer_fold}") + + train_result = train_fixed_epochs( + model=model, + train_loader=train_loader, + val_loader=monitor_loader, + device=device, + lr=best_params["lr"], + weight_decay=best_params["weight_decay"], + epochs=epoch_mean, + class_weights=class_weights, + use_cosine_annealing=True, + backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0), + freeze_backbone_epochs=3, + ) + + model.load_state_dict(train_result["final_state"]) + model = model.to(device) + + config = { + "d_model": best_params["d_model"], + "num_heads": best_params["num_heads"], + "n_attn_layers": best_params["n_attn_layers"], + "fusion_strategy": best_params["fusion_strategy"], + "head_hidden_dim": best_params["head_hidden_dim"], + "set_transformer_block": best_params.get("set_transformer_block", "sab"), + "dropout": best_params["dropout"], + "use_mpnn": use_mpnn, + "use_moe": use_moe, + "moe_n_experts": moe_n_experts, + "moe_top_k": moe_top_k, + "moe_expert_hidden_mult": moe_expert_hidden_mult, + "moe_jitter_noise": moe_jitter_noise, + **(llm_kwargs or {}), + } + + torch.save({ + "model_state_dict": train_result["final_state"], + "config": config, + "epoch_mean": epoch_mean, + "best_params": best_params, + }, fold_dir / "model.pt") + + with open(fold_dir / "history.json", "w") as f: + json.dump(train_result["history"], f, indent=2) + + # 在测试集上评估 + logger.info("Evaluating on outer test set...") + test_metrics = evaluate_on_test(model, test_loader, device) + + with open(fold_dir / "test_metrics.json", "w") as f: + json.dump(test_metrics, f, indent=2) + + logger.info(f"\nOuter Fold {outer_fold} Test Results:") + for task, metrics in test_metrics.items(): + if "rmse" in metrics: + logger.info(f" {task}: RMSE={metrics['rmse']:.4f}, R²={metrics['r2']:.4f}") + elif "accuracy" in metrics: + logger.info(f" {task}: Acc={metrics['accuracy']:.4f}, F1={metrics['f1']:.4f}") + elif "kl_divergence" in metrics: + logger.info(f" {task}: KL={metrics['kl_divergence']:.4f}, JS={metrics['js_divergence']:.4f}") + + return { + "fold": outer_fold, + "best_params": best_params, + "epoch_mean": epoch_mean, + "test_metrics": test_metrics, + } + + +# ============ 主流程 ============ +@app.command() +def main( + input_path: Path = INTERIM_DATA_DIR / "internal.csv", + output_dir: Path = MODELS_DIR / "nested_cv", + # CV 参数 + n_outer_folds: int = 5, + n_inner_folds: int = 3, + min_stratum_count: int = 5, + seed: int = 42, + # Optuna 参数 + n_trials: int = 20, + epochs_per_trial: int = 30, + inner_patience: int = 10, + # 训练参数 + batch_size: int = 32, + # 预训练权重 + init_from_pretrain: Optional[Path] = None, + load_delivery_head: bool = False, + # MPNN + use_mpnn: bool = False, + n_repeats: int = 1, + repeat_seed_step: int = 1000, + # MoE(消融开关) + use_moe: bool = False, + moe_n_experts: int = 4, + moe_top_k: int = 2, + moe_expert_hidden_mult: int = 2, + moe_jitter_noise: float = 0.0, + # Set Transformer + set_transformer_block: str = "sab", + # LLM(消融开关) + use_llm: bool = False, + llm_model_path: str = DEFAULT_MOLT5_PATH, + llm_freeze: bool = True, + llm_use_lora: bool = False, + llm_lora_r: int = 8, + llm_lora_alpha: int = 16, + llm_lora_dropout: float = 0.05, + # 并行 + parallel: bool = False, + # 设备 + device: str = "cuda" if torch.cuda.is_available() else "cpu", +): + """ + 嵌套交叉验证 + Optuna 超参调优。 + + 外层 5-fold(20% test / 80% train),内层 3-fold Optuna 调参。 + 外层训练不使用 early-stopping,epoch 数使用内层 best trial 的 epoch_mean。 + + 使用 --init-from-pretrain 从预训练 checkpoint 初始化模型权重。 + 使用 --parallel 同时运行所有外层 fold(需要足够 GPU 显存)。 + """ + if optuna is None: + logger.error("Optuna not installed. Run: pip install optuna") + raise typer.Exit(1) + + logger.info(f"Using device: {device}") + device = torch.device(device) + set_global_seed(seed) + + llm_kwargs = dict( + use_llm=use_llm, + llm_model_path=llm_model_path, + llm_freeze=llm_freeze, + llm_use_lora=llm_use_lora, + llm_lora_r=llm_lora_r, + llm_lora_alpha=llm_lora_alpha, + llm_lora_dropout=llm_lora_dropout, + ) + + # 加载预训练权重(如果指定) + pretrain_state_dict = None + pretrain_config = None + if init_from_pretrain is not None: + if init_from_pretrain.exists(): + logger.info(f"Loading pretrain weights from {init_from_pretrain}") + checkpoint = torch.load(init_from_pretrain, map_location="cpu", weights_only=False) + pretrain_state_dict = checkpoint["model_state_dict"] + pretrain_config = checkpoint.get("config", {}) + logger.success(f"Loaded pretrain checkpoint (d_model={pretrain_config.get('d_model')})") + else: + logger.warning(f"Pretrain checkpoint not found: {init_from_pretrain}, skipping") + + # 创建输出目录(带时间戳) + run_name = datetime.now().strftime("%Y%m%d_%H%M%S") + run_dir = output_dir / run_name + run_dir.mkdir(parents=True, exist_ok=True) + logger.info(f"Output directory: {run_dir}") + + # 加载数据 + logger.info(f"Loading data from {input_path}") + df = pd.read_csv(input_path) + logger.info(f"Loaded {len(df)} samples") + + # 处理数据 + logger.info("Processing dataframe...") + df = process_dataframe(df) + + # 构建复合分层标签 + logger.info("Building composite strata...") + strata, strata_info = build_composite_strata(df, min_stratum_count) + + # 保存 strata 信息 + with open(run_dir / "strata_info.json", "w") as f: + json.dump(strata_info, f, indent=2, default=str) + + # 创建完整数据集(仅用于获取样本数做 split) + n_samples = len(LNPDataset(df)) + + # 外层 CV split + outer_cv = StratifiedKFold( + n_splits=n_outer_folds, shuffle=True, random_state=seed + ) + + device_str = str(device) + fold_args = [] + for outer_fold, (outer_train_idx, outer_test_idx) in enumerate( + outer_cv.split(np.arange(n_samples), strata) + ): + fold_args.append(dict( + outer_fold=outer_fold, + outer_train_idx=outer_train_idx, + outer_test_idx=outer_test_idx, + df=df, + strata=strata, + fold_dir=run_dir / f"outer_fold_{outer_fold}", + n_trials=n_trials, + epochs_per_trial=epochs_per_trial, + inner_patience=inner_patience, + batch_size=batch_size, + n_inner_folds=n_inner_folds, + use_mpnn=use_mpnn, + seed=seed, + pretrain_state_dict=pretrain_state_dict, + 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, + set_transformer_block=set_transformer_block, + llm_kwargs=llm_kwargs, + )) + + if parallel: + import multiprocessing as mp + from concurrent.futures import ProcessPoolExecutor + + ctx = mp.get_context("spawn") + logger.info(f"Running {n_outer_folds} outer folds in PARALLEL (spawn)") + with ProcessPoolExecutor(max_workers=n_outer_folds, mp_context=ctx) as executor: + futures = [executor.submit(_run_single_outer_fold, **args) for args in fold_args] + outer_results = [f.result() for f in futures] + outer_results.sort(key=lambda r: r["fold"]) + else: + logger.info(f"Running {n_outer_folds} outer folds SEQUENTIALLY") + outer_results = [] + per_fold_params = {} + for args in fold_args: + result = _run_single_outer_fold(**args) + outer_results.append(result) + per_fold_params[result["fold"]] = (result["best_params"], result["epoch_mean"]) + + for rep in range(1, n_repeats): + logger.info(f"\n===== Repeat {rep}/{n_repeats - 1} (reuse hyperparams) =====") + for args in fold_args: + bp, em = per_fold_params[args["outer_fold"]] + rep_args = dict(args) + rep_args["seed"] = seed + rep * repeat_seed_step + rep_args["fold_dir"] = run_dir / f"repeat_{rep}" / f"outer_fold_{args['outer_fold']}" + rep_args["precomputed_best_params"] = bp + rep_args["precomputed_epoch_mean"] = em + outer_results.append(_run_single_outer_fold(**rep_args)) + + # 汇总结果 + logger.info("\n" + "=" * 60) + logger.info("NESTED CV COMPLETE") + logger.info("=" * 60) + + # 计算汇总统计 + summary = {"fold_results": outer_results} + + # 对每个任务计算均值和标准差 + tasks_with_metrics = {} + for result in outer_results: + for task, metrics in result["test_metrics"].items(): + if task not in tasks_with_metrics: + tasks_with_metrics[task] = {k: [] for k in metrics.keys() if k != "n_samples"} + for k, v in metrics.items(): + if k != "n_samples": + tasks_with_metrics[task][k].append(v) + + summary["summary_stats"] = {} + for task, metrics_dict in tasks_with_metrics.items(): + summary["summary_stats"][task] = {} + for metric_name, values in metrics_dict.items(): + summary["summary_stats"][task][f"{metric_name}_mean"] = float(np.mean(values)) + summary["summary_stats"][task][f"{metric_name}_std"] = float(np.std(values)) + + # 打印汇总 + logger.info("\n[Summary Statistics]") + for task, stats in summary["summary_stats"].items(): + if "rmse_mean" in stats: + logger.info( + f" {task}: RMSE={stats['rmse_mean']:.4f}±{stats['rmse_std']:.4f}, " + f"R²={stats['r2_mean']:.4f}±{stats['r2_std']:.4f}" + ) + elif "accuracy_mean" in stats: + logger.info( + f" {task}: Acc={stats['accuracy_mean']:.4f}±{stats['accuracy_std']:.4f}, " + f"F1={stats['f1_mean']:.4f}±{stats['f1_std']:.4f}" + ) + elif "kl_divergence_mean" in stats: + logger.info( + f" {task}: KL={stats['kl_divergence_mean']:.4f}±{stats['kl_divergence_std']:.4f}, " + f"JS={stats['js_divergence_mean']:.4f}±{stats['js_divergence_std']:.4f}" + ) + + # 保存汇总 + with open(run_dir / "summary.json", "w") as f: + json.dump(summary, f, indent=2) + + logger.success(f"\nAll results saved to {run_dir}") + + +if __name__ == "__main__": + app() + diff --git a/lnp_ml/modeling/predict.py b/lnp_ml/modeling/predict.py index 6722da8..a5cf836 100644 --- a/lnp_ml/modeling/predict.py +++ b/lnp_ml/modeling/predict.py @@ -55,6 +55,13 @@ def load_model( fusion_strategy=config["fusion_strategy"], head_hidden_dim=config["head_hidden_dim"], dropout=config["dropout"], + use_llm=config.get("use_llm", False), + llm_model_path=config.get("llm_model_path", "models/molt5-base"), + llm_freeze=config.get("llm_freeze", True), + llm_use_lora=config.get("llm_use_lora", False), + llm_lora_r=config.get("llm_lora_r", 8), + llm_lora_alpha=config.get("llm_lora_alpha", 16), + llm_lora_dropout=config.get("llm_lora_dropout", 0.05), mpnn_ensemble_paths=ensemble_paths, mpnn_device=mpnn_device, ) @@ -66,9 +73,16 @@ def load_model( fusion_strategy=config["fusion_strategy"], head_hidden_dim=config["head_hidden_dim"], dropout=config["dropout"], + use_llm=config.get("use_llm", False), + llm_model_path=config.get("llm_model_path", "models/molt5-base"), + llm_freeze=config.get("llm_freeze", True), + llm_use_lora=config.get("llm_use_lora", False), + llm_lora_r=config.get("llm_lora_r", 8), + llm_lora_alpha=config.get("llm_lora_alpha", 16), + llm_lora_dropout=config.get("llm_lora_dropout", 0.05), ) - model.load_state_dict(checkpoint["model_state_dict"]) + model.load_state_dict(checkpoint["model_state_dict"], strict=False) model.to(device) model.eval()