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:
DicongLi 2026-07-15 14:21:23 +08:00
parent 0c6828076b
commit 104dfef94c
4 changed files with 143 additions and 128 deletions

View File

@ -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),
}

View File

@ -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

View File

@ -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,

View File

@ -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,