TW-DTI / eval.py
graenys's picture
Upload 9 files
4e2dbd1 verified
Raw
History Blame Contribute Delete
5.5 kB
import torch
from torch.utils.data import DataLoader
from transformers import AutoTokenizer, EsmTokenizer
from tqdm import tqdm
import csv
import numpy as np
from sklearn.metrics import accuracy_score, recall_score, precision_score, f1_score, roc_auc_score, confusion_matrix, classification_report
# 引入你的自定义模块
from daul_tower_model.net import DualTowerModel
from daul_tower_model.config import Config
from daul_tower_model.dataset import DrugTargetDataset, DualTowerCollator
# === 配置部分 ===
# 模型权重路径 (修改为你想要测试的那个 epoch)
CHECKPOINT_PATH = "daul_checkpoints/model_epoch_100.pth"
# 测试数据路径 (修改为你要测试的 CSV 文件)
TEST_CSV_PATH = "daul_tower_model/data/drug_target_activity_sequence_compound_smiles_nnew.csv"
BATCH_SIZE = 256 # 推理时不存梯度,Batch Size 可以比训练时大一倍
def load_test_data(csv_path):
"""读取 CSV 并转换为 Dataset 需要的 List 格式"""
data_list = []
print(f"Reading test data from {csv_path}...")
with open(csv_path, "r", encoding="utf-8") as f:
reader = csv.DictReader(f)
for row in reader:
prot_seq = row.get("sequence", "").strip()
mol_smiles = row.get("compound__smiles", "").strip()
# 兼容 outcome_is_active 列名,或者 label 列名
raw_label = row.get("outcome_is_active", row.get("label", "")).strip()
if prot_seq and mol_smiles and raw_label:
# 简单的 Label 转换逻辑
if raw_label.lower() in ['true', '1', '1.0']:
label = 1.0
else:
label = 0.0
data_list.append((prot_seq, mol_smiles, label))
print(f"Loaded {len(data_list)} test samples.")
return data_list
def main():
# 1. 设备配置
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Running evaluation on: {device}")
# 2. 加载 Tokenizers
print("Loading Tokenizers...")
prot_tokenizer = EsmTokenizer.from_pretrained(Config.SAPROT_PATH)
mol_tokenizer = AutoTokenizer.from_pretrained(Config.CHEMBERTA_PATH)
# 3. 准备数据
test_data = load_test_data(TEST_CSV_PATH)
# 注意:测试时不需要 DynamicBalancedDataset,直接用最普通的 Dataset 即可,我们要测真实分布
# 如果你的 Dataset 类里还有随机采样逻辑,请务必使用下面这个简单的临时类,或者确保你的 Dataset 没做任何 Drop
dataset = DrugTargetDataset(test_data)
collator = DualTowerCollator(prot_tokenizer, mol_tokenizer)
dataloader = DataLoader(
dataset,
batch_size=BATCH_SIZE,
shuffle=False, # 测试不需要打乱
collate_fn=collator,
num_workers=4
)
# 4. 加载模型
print("Loading Model...")
model = DualTowerModel(Config)
# 加载权重
print(f"Loading weights from {CHECKPOINT_PATH}...")
state_dict = torch.load(CHECKPOINT_PATH, map_location=device)
model.load_state_dict(state_dict)
model.to(device)
model.eval() # ⚠️ 关键:切换到评估模式 (关闭 Dropout 等)
# 5. 开始推理 (Inference Loop)
all_targets = []
all_preds = []
all_probs = []
print("Starting Inference...")
with torch.no_grad(): # ⚠️ 关键:关闭梯度计算,节省显存
for prot_inputs, mol_inputs, labels in tqdm(dataloader):
# 移到 GPU
prot_inputs = {k: v.to(device) for k, v in prot_inputs.items()}
mol_inputs = {k: v.to(device) for k, v in mol_inputs.items()}
labels = labels.to(device)
# 前向传播
# forward 返回 (prot_feat, mol_feat, similarity_logits)
_, _, logits = model(prot_inputs, mol_inputs)
# 维度对齐
if len(logits.shape) > 1:
logits = logits.view(-1)
# 计算概率 (Sigmoid)
probs = torch.sigmoid(logits)
# 存回 CPU 以便用 sklearn 计算指标
all_probs.extend(probs.cpu().numpy())
all_targets.extend(labels.cpu().numpy())
# 6. 计算指标
all_targets = np.array(all_targets)
all_probs = np.array(all_probs)
# 阈值设为 0.5,大于 0.5 算正样本
all_preds = (all_probs > 0.5).astype(int)
print("\n" + "="*30)
print("📊 Evaluation Results")
print("="*30)
# 基础指标
acc = accuracy_score(all_targets, all_preds)
rec = recall_score(all_targets, all_preds)
prec = precision_score(all_targets, all_preds, zero_division=0)
f1 = f1_score(all_targets, all_preds, zero_division=0)
try:
auc = roc_auc_score(all_targets, all_probs)
except:
auc = 0.0
print("Warning: AUC calculation failed (possibly only one class present).")
print(f"Accuracy : {acc:.2%}")
print(f"Precision : {prec:.2%}")
print(f"Recall : {rec:.2%} <-- 最关键指标")
print(f"F1 Score : {f1:.2%}")
print(f"ROC AUC : {auc:.4f}")
print("\n[Confusion Matrix]")
cm = confusion_matrix(all_targets, all_preds)
print(f"TN: {cm[0][0]} | FP: {cm[0][1]}")
print(f"FN: {cm[1][0]} | TP: {cm[1][1]}")
print("\n[Detailed Report]")
print(classification_report(all_targets, all_preds, target_names=["Inactive (0)", "Active (1)"]))
if __name__ == "__main__":
main()