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()