mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 20:33:18 +08:00
129 lines
4.2 KiB
Python
129 lines
4.2 KiB
Python
"""LNPModel 整图自检脚本
|
||
|
||
base/+moe 路径只需 rdkit + torch;+llm 路径需 MolT5。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import sys
|
||
import traceback
|
||
from typing import Callable, Dict, List, Tuple
|
||
|
||
import torch
|
||
|
||
from lnp_ml.featurization.smiles import RDKitFeaturizer
|
||
from lnp_ml.modeling.models import LNPModelWithoutMPNN
|
||
|
||
MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
|
||
|
||
# 期望的任务输出维度
|
||
EXPECTED = {"size": 1, "pdi": 2, "ee": 3, "delivery": 1, "biodist": 7, "toxic": 2}
|
||
|
||
|
||
def _input_dims() -> Dict[str, int]:
|
||
"""根据实际 RDKit 输出确定无 MPNN 时的输入维度。"""
|
||
feat = RDKitFeaturizer().transform(["CCO"])
|
||
return {
|
||
"morgan": feat["morgan"].shape[1],
|
||
"maccs": feat["maccs"].shape[1],
|
||
"desc": feat["desc"].shape[1],
|
||
"comp": 5, "phys": 12, "help": 4, "exp": 32,
|
||
}
|
||
|
||
|
||
def _batch(B: int = 2) -> Tuple[List[str], Dict[str, torch.Tensor]]:
|
||
smiles = ["CCO", "CCN"][:B]
|
||
tabular = {
|
||
"comp": torch.randn(B, 5),
|
||
"phys": torch.randn(B, 12),
|
||
"help": torch.randn(B, 4),
|
||
"exp": torch.randn(B, 32),
|
||
}
|
||
return smiles, tabular
|
||
|
||
|
||
def _check_outputs(outputs: Dict[str, torch.Tensor], B: int) -> None:
|
||
for task, dim in EXPECTED.items():
|
||
assert outputs[task].shape == (B, dim), f"{task}: {outputs[task].shape} != ({B}, {dim})"
|
||
|
||
|
||
def test_forward_base() -> None:
|
||
model = LNPModelWithoutMPNN(d_model=64, input_dims=_input_dims()).eval()
|
||
smiles, tabular = _batch()
|
||
_check_outputs(model(smiles, tabular), B=2)
|
||
|
||
|
||
def test_forward_moe() -> None:
|
||
model = LNPModelWithoutMPNN(d_model=64, input_dims=_input_dims(), use_moe=True).eval()
|
||
smiles, tabular = _batch()
|
||
_check_outputs(model(smiles, tabular), B=2)
|
||
assert model.get_last_moe_extras() is not None, "use_moe=True 应产生 MoE extras"
|
||
|
||
|
||
def test_forward_from_projected() -> None:
|
||
model = LNPModelWithoutMPNN(d_model=64, input_dims=_input_dims()).eval()
|
||
smiles, tabular = _batch()
|
||
stacked = model._encode_and_project(smiles, tabular)
|
||
out = model.forward_from_projected(stacked, task="biodist")
|
||
assert out.shape == (2, 7), f"biodist {out.shape} != (2, 7)"
|
||
|
||
|
||
def test_split_idx_no_mpnn() -> None:
|
||
model = LNPModelWithoutMPNN(d_model=64, input_dims=_input_dims())
|
||
assert model.split_idx == 3, "无 MPNN 时化学 token 应为 3"
|
||
|
||
|
||
def test_forward_llm() -> None:
|
||
"""需要 MolT5;缺失则跳过。"""
|
||
try:
|
||
model = LNPModelWithoutMPNN(
|
||
d_model=64, input_dims=_input_dims(),
|
||
use_moe=True, use_llm=True, llm_model_path=MOLT5_PATH,
|
||
).eval()
|
||
except Exception as e:
|
||
print(f"[SKIP] test_forward_llm(MolT5 不可用):{e}")
|
||
return
|
||
smiles, tabular = _batch()
|
||
_check_outputs(model(smiles, tabular), B=2)
|
||
# 冻结的 MolT5 encoder 不应进入 backbone state dict
|
||
bb = model.get_backbone_state_dict()
|
||
assert not any(k.startswith("llm_prompt.encoder.") for k in bb), "backbone 不应含冻结 encoder 权重"
|
||
|
||
|
||
# ──────────────────────────────────────────────────────────────────────
|
||
# Runner
|
||
# ──────────────────────────────────────────────────────────────────────
|
||
|
||
def _collect_tests() -> List[Tuple[str, Callable[[], None]]]:
|
||
return [
|
||
(name, fn) for name, fn in globals().items()
|
||
if name.startswith("test_") and callable(fn)
|
||
]
|
||
|
||
|
||
def main() -> int:
|
||
tests = _collect_tests()
|
||
pass_count = 0
|
||
fail: List[Tuple[str, str]] = []
|
||
for name, fn in tests:
|
||
try:
|
||
fn()
|
||
except Exception:
|
||
fail.append((name, traceback.format_exc()))
|
||
print(f"[FAIL] {name}")
|
||
else:
|
||
pass_count += 1
|
||
print(f"[ OK ] {name}")
|
||
print(f"\n{'='*60}")
|
||
print(f"Passed: {pass_count}/{len(tests)}")
|
||
if fail:
|
||
print(f"\n{'='*60}\nFailures:")
|
||
for name, tb in fail:
|
||
print(f"\n--- {name} ---\n{tb}")
|
||
return 1
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main()) |