| 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 |
|
|
| |
| |
| CHECKPOINT_PATH = "daul_checkpoints/model_epoch_100.pth" |
| |
| TEST_CSV_PATH = "daul_tower_model/data/drug_target_activity_sequence_compound_smiles_nnew.csv" |
| BATCH_SIZE = 256 |
|
|
| 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() |
| |
| raw_label = row.get("outcome_is_active", row.get("label", "")).strip() |
|
|
| if prot_seq and mol_smiles and raw_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(): |
| |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| print(f"Running evaluation on: {device}") |
|
|
| |
| print("Loading Tokenizers...") |
| prot_tokenizer = EsmTokenizer.from_pretrained(Config.SAPROT_PATH) |
| mol_tokenizer = AutoTokenizer.from_pretrained(Config.CHEMBERTA_PATH) |
|
|
| |
| test_data = load_test_data(TEST_CSV_PATH) |
| |
| |
| 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 |
| ) |
|
|
| |
| 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() |
|
|
| |
| all_targets = [] |
| all_preds = [] |
| all_probs = [] |
|
|
| print("Starting Inference...") |
| with torch.no_grad(): |
| for prot_inputs, mol_inputs, labels in tqdm(dataloader): |
| |
| 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) |
|
|
| |
| |
| _, _, logits = model(prot_inputs, mol_inputs) |
|
|
| |
| if len(logits.shape) > 1: |
| logits = logits.view(-1) |
| |
| |
| probs = torch.sigmoid(logits) |
| |
| |
| all_probs.extend(probs.cpu().numpy()) |
| all_targets.extend(labels.cpu().numpy()) |
|
|
| |
| all_targets = np.array(all_targets) |
| all_probs = np.array(all_probs) |
| |
| 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() |