mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-21 21:22:05 +08:00
feat: 回归旁路(regression bypass)+ --reg-bypass 开关
- fusion.py: ResidualConcatFusion 额外返回 pooled(纯 chem+tab 注意力池化) - models.py: _backbone_from_stacked 存 _last_pooled;forward 按 reg_bypass 决定是否传入 - heads.py: MultiTaskHead 接收 x_numeric,回归头(size/delivery)走纯数值,分类/biodist 走 fused - nested_cv_optuna.py: 新增 --reg-bypass 开关(typer),经 llm_kwargs 传入 fix: LNPModelWithoutMPNN 未将 reg_bypass 转发给 super().__init__(), 导致开关静默失效、on/off 产生完全相同结果(已验证修复:off→False, on→True)
This commit is contained in:
parent
0c6828076b
commit
104dfef94c
@ -101,7 +101,7 @@ class MultiTaskHead(nn.Module):
|
||||
for t in ["size", "delivery", "pdi", "ee", "toxic", "biodist"]
|
||||
})
|
||||
|
||||
def forward(self, x: torch.Tensor) -> Dict[str, torch.Tensor]:
|
||||
def forward(self, x: torch.Tensor, x_numeric: torch.Tensor = None) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Args:
|
||||
x: [B, in_dim] fusion 层输出
|
||||
@ -115,11 +115,12 @@ class MultiTaskHead(nn.Module):
|
||||
- "biodist": [B, 7] probabilities (sum=1)
|
||||
- "toxic": [B, 2] logits
|
||||
"""
|
||||
xn = x if x_numeric is None else x_numeric
|
||||
return {
|
||||
"size": self.size_head(x),
|
||||
"size": self.size_head(xn),
|
||||
"pdi": self.pdi_head(x),
|
||||
"ee": self.ee_head(x),
|
||||
"delivery": self.delivery_head(x),
|
||||
"delivery": self.delivery_head(xn),
|
||||
"biodist": self.biodist_head(x),
|
||||
"toxic": self.toxic_head(x),
|
||||
}
|
||||
|
||||
@ -153,4 +153,7 @@ class ResidualConcatFusion(nn.Module):
|
||||
if f_retr is not None:
|
||||
out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开
|
||||
|
||||
return (out, attn) if return_attn_weights else out
|
||||
# 回归旁路:额外返回 pooled(纯 chem+tab,不含 f_llm/f_moe/f_retr)
|
||||
if return_attn_weights:
|
||||
return out, pooled, attn
|
||||
return out, pooled
|
||||
@ -99,6 +99,8 @@ class LNPModel(nn.Module):
|
||||
moe_top_k: int = 2,
|
||||
moe_expert_hidden_mult: int = 2,
|
||||
moe_jitter_noise: float = 0.0,
|
||||
# ============ 回归旁路开关 ============
|
||||
reg_bypass: str = "on",
|
||||
# ============ LLM 相关 ============
|
||||
use_llm: bool = False,
|
||||
llm_model_path: str = DEFAULT_MOLT5_PATH,
|
||||
@ -174,6 +176,7 @@ class LNPModel(nn.Module):
|
||||
|
||||
# ============ LLM Prompt (可选) ============
|
||||
self.use_llm = use_llm
|
||||
self.reg_bypass = (str(reg_bypass).lower() == "on")
|
||||
if use_llm:
|
||||
self.llm_prompt = LLMPromptEncoder(
|
||||
d_model=d_model,
|
||||
@ -270,7 +273,9 @@ class LNPModel(nn.Module):
|
||||
_feats_t = torch.as_tensor(_feats, dtype=chem.dtype, device=chem.device)
|
||||
f_retr = self.retr_proj(_feats_t)
|
||||
|
||||
return self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm, f_retr=f_retr)
|
||||
fused, pooled = self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm, f_retr=f_retr)
|
||||
self._last_pooled = pooled # 纯数值向量,供回归 head 使用
|
||||
return fused
|
||||
|
||||
def forward_from_projected(
|
||||
self,
|
||||
@ -357,7 +362,7 @@ class LNPModel(nn.Module):
|
||||
delivery [B,1], biodist [B,7], toxic [B,2]
|
||||
"""
|
||||
fused = self.forward_backbone(smiles, tabular)
|
||||
return self.head(fused)
|
||||
return self.head(fused, getattr(self, "_last_pooled", None) if getattr(self, "reg_bypass", True) else None)
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
"""清空所有 encoder 的缓存"""
|
||||
@ -443,6 +448,8 @@ class LNPModelWithoutMPNN(LNPModel):
|
||||
moe_top_k: int = 2,
|
||||
moe_expert_hidden_mult: int = 2,
|
||||
moe_jitter_noise: float = 0.0,
|
||||
# ============ 回归旁路开关 ============
|
||||
reg_bypass: str = "on",
|
||||
# ============ LLM 相关 ============
|
||||
use_llm: bool = False,
|
||||
llm_model_path: str = DEFAULT_MOLT5_PATH,
|
||||
@ -472,6 +479,7 @@ class LNPModelWithoutMPNN(LNPModel):
|
||||
mpnn_checkpoint=None,
|
||||
mpnn_ensemble_paths=None,
|
||||
input_dims=dims,
|
||||
reg_bypass=reg_bypass,
|
||||
use_moe=use_moe,
|
||||
moe_n_experts=moe_n_experts,
|
||||
moe_top_k=moe_top_k,
|
||||
|
||||
@ -930,6 +930,8 @@ def main(
|
||||
moe_jitter_noise: float = 0.0,
|
||||
# Set Transformer
|
||||
set_transformer_block: str = "sab",
|
||||
# 回归旁路(消融开关)
|
||||
reg_bypass: str = "on",
|
||||
# LLM(消融开关)
|
||||
use_llm: bool = False,
|
||||
use_rag: bool = False,
|
||||
@ -967,6 +969,7 @@ def main(
|
||||
set_global_seed(seed)
|
||||
|
||||
llm_kwargs = dict(
|
||||
reg_bypass=reg_bypass,
|
||||
use_rag=use_rag,
|
||||
rag_top_k=rag_top_k,
|
||||
use_llm=use_llm,
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user