AIVS / REDDA /main.py
yg3191's picture
Upload folder using huggingface_hub
132149b verified
Raw
History Blame Contribute Delete
12.6 kB
"""
Module for Drug Repositioning (Drug-Disease Association Prediction).
This module implements the training pipeline for a drug-disease association (DDA) prediction model.
It performs K-fold cross-validation on the dataset, constructs training/testing splits, and trains the model
with early stopping. The model predictions are then saved and evaluation plots (ROC and PR curves) are generated.
"""
import os
import numpy as np
import pandas as pd
import torch as th
from warnings import simplefilter
from sklearn.model_selection import KFold
import dgl
import time
from model import Model
from load_data import load, remove_graph
from utils import (
get_metrics_auc,
set_seed,
plot_result_auc,
plot_result_aupr,
EarlyStopping,
get_metrics,
)
from args import args
def run_training():
"""
Train the drug-disease association prediction model using K-fold cross-validation.
Full-graph training (no k-hop subgraph), and for KG datasets perform training-time
negative sampling (20x positives) while leaving the test set unchanged.
"""
# Ignore FutureWarning messages
simplefilter(action="ignore", category=FutureWarning)
# Print arguments for debugging
print(args)
set_seed(args.seed)
# Create the directory to save results if it doesn't exist
os.makedirs(args.saved_path, exist_ok=True)
# Select device: GPU if device_id is provided, else CPU
if args.device_id:
print("Training on GPU")
device = th.device(f"cuda:{args.device_id}")
else:
print("Training on CPU")
device = th.device("cpu")
# Build DDA matrix
if args.dataset in ("KGdataset", "KGdataset_tiny"):
print("Processing KGdataset...")
g_tmp = load(args.dataset)
num_drugs = g_tmp.num_nodes("drug")
num_diseases = g_tmp.num_nodes("disease")
drug_disease_matrix = np.zeros((num_drugs, num_diseases), dtype=np.int64)
if "drug_disease" in g_tmp.etypes:
drug_disease_edges = g_tmp.edges(etype="drug_disease")
di = drug_disease_edges[0].cpu().numpy()
dj = drug_disease_edges[1].cpu().numpy()
drug_disease_matrix[di, dj] = 1
print(f"Created drug-disease matrix: {num_drugs} drugs x {num_diseases} diseases")
print(f"Number of positive associations: {int(drug_disease_matrix.sum())}")
df = drug_disease_matrix
else:
df = pd.read_csv(
f"./dataset/{args.dataset}/{args.dataset}_baseline.csv", header=None
).values
# Construct an array with rows: [drug_index, disease_index, label]
data = np.array([[i, j, df[i, j]] for i in range(df.shape[0]) for j in range(df.shape[1])], dtype="int64")
data_pos = data[data[:, -1] == 1]
data_neg = data[data[:, -1] == 0]
assert len(data) == len(data_pos) + len(data_neg)
print(f"Total samples: {len(data)}")
print(f"Positive samples: {len(data_pos)}")
print(f"Negative samples: {len(data_neg)}")
# Set random seed and initialize KFold cross-validator
set_seed(args.seed)
kf = KFold(n_splits=args.nfold, shuffle=True, random_state=args.seed)
fold = 1
pred_result = np.zeros(df.shape)
# Loop over K-fold splits for positive and negative samples simultaneously
for (train_pos_idx, test_pos_idx), (train_neg_idx, test_neg_idx) in zip(
kf.split(data_pos), kf.split(data_neg)
):
print(f"{args.nfold}-Fold Cross Validation: Fold {fold}")
# Get training and testing indices for positive and negative samples
train_pos_id, test_pos_id = data_pos[train_pos_idx], data_pos[test_pos_idx]
train_neg_id, test_neg_id = data_neg[train_neg_idx], data_neg[test_neg_idx]
# check if train_pos_id and test_pos_id are out of bound of df
assert max(train_pos_id[:, 0]) < df.shape[0] and max(test_pos_id[:, 0]) < df.shape[0]
assert max(train_pos_id[:, 1]) < df.shape[1] and max(test_pos_id[:, 1]) < df.shape[1]
assert max(train_neg_id[:, 0]) < df.shape[0] and max(test_neg_id[:, 0]) < df.shape[0]
assert max(train_neg_id[:, 1]) < df.shape[1] and max(test_neg_id[:, 1]) < df.shape[1]
# Create index tuples for easier tensor indexing later
train_pos_indices = [tuple(train_pos_id[:, 0]), tuple(train_pos_id[:, 1])]
test_pos_indices = [tuple(test_pos_id[:, 0]), tuple(test_pos_id[:, 1])]
train_neg_indices = [tuple(train_neg_id[:, 0]), tuple(train_neg_id[:, 1])]
test_neg_indices = [tuple(test_neg_id[:, 0]), tuple(test_neg_id[:, 1])]
assert (
len(train_pos_indices[0])
+ len(test_pos_indices[0])
+ len(train_neg_indices[0])
+ len(test_neg_indices[0])
== len(data)
)
# Load graph data and remove test positives from the training graph
g = load(args.dataset)
print(g)
g = remove_graph(g, test_pos_id[:, :-1]).to(device)
# Extract node features
if args.dataset == "Kdataset":
feature = {
"drug": g.nodes["drug"].data["h"],
"disease": g.nodes["disease"].data["h"],
"protein": g.nodes["protein"].data["h"],
"gene": g.nodes["gene"].data["h"],
"pathway": g.nodes["pathway"].data["h"],
}
elif args.dataset == "Bdataset":
feature = {
"drug": g.nodes["drug"].data["h"],
"disease": g.nodes["disease"].data["h"],
"protein": g.nodes["protein"].data["h"],
}
elif args.dataset in ("KGdataset", "KGdataset_tiny"):
feature = {ntype: g.nodes[ntype].data["h"] for ntype in g.ntypes}
print(f"Node types in graph: {list(feature.keys())}")
else:
raise ValueError("Unsupported dataset type.")
# Create mask labels: set test positions to 0, others to 1
mask_label = np.ones(df.shape, dtype=np.int8)
mask_label[test_pos_indices[0], test_pos_indices[1]] = 0
mask_label[test_neg_indices[0], test_neg_indices[1]] = 0
# Test mask (no sampling)
mask_test_np = np.where(mask_label == 0)
mask_test = [tuple(mask_test_np[0]), tuple(mask_test_np[1])]
# Training mask
if args.dataset in ("KGdataset", "KGdataset_tiny"):
# Sample negatives at 20x positives (training only)
n_pos = len(train_pos_indices[0])
n_neg_total = len(train_neg_indices[0])
n_neg_sample = min(20 * n_pos, n_neg_total)
rng = np.random.default_rng(args.seed + fold)
sel = rng.choice(n_neg_total, size=n_neg_sample, replace=False)
neg0 = np.array(train_neg_indices[0])[sel]
neg1 = np.array(train_neg_indices[1])[sel]
tr0 = np.concatenate([np.array(train_pos_indices[0]), neg0])
tr1 = np.concatenate([np.array(train_pos_indices[1]), neg1])
mask_train = [tuple(tr0), tuple(tr1)]
pos_weight = th.tensor(n_neg_sample / max(1, n_pos))
else:
mask_train_np = np.where(mask_label == 1)
mask_train = [tuple(mask_train_np[0]), tuple(mask_train_np[1])]
pos_weight = th.tensor(len(train_neg_indices[0]) / max(1, len(train_pos_indices[0])))
print(
"Number of total training samples: {}, pos samples: {}, neg samples (used): {}".format(
len(mask_train[0]),
len(train_pos_indices[0]),
len(mask_train[0]) - len(train_pos_indices[0]),
)
)
print(
"Number of total testing samples: {}, pos samples: {}, neg samples: {}".format(
len(mask_test[0]),
len(test_pos_indices[0]),
len(test_neg_indices[0]),
)
)
label = th.tensor(df).float().to(device)
# Initialize the model and optimizer
model = Model(
etypes=g.etypes,
ntypes=g.ntypes,
in_feats=feature["drug"].shape[1],
hidden_feats=args.hidden_feats,
num_heads=args.num_heads,
dropout=args.dropout,
).to(device)
optimizer = th.optim.Adam(
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
)
optim_scheduler = th.optim.lr_scheduler.CyclicLR(
optimizer,
base_lr=0.1 * args.learning_rate,
max_lr=args.learning_rate,
gamma=0.995,
step_size_up=20,
mode="exp_range",
cycle_momentum=False,
)
criterion = th.nn.BCEWithLogitsLoss(pos_weight=pos_weight)
print(f"Loss pos weight: {pos_weight.item():.3f}")
stopper = EarlyStopping(patience=args.patience, saved_path=args.saved_path)
# Full-graph training (no subgraph)
for epoch in range(1, args.epoch + 1):
time_start = time.time()
model.train()
score = model(g, feature)
pred = th.sigmoid(score)
loss = criterion(
score[mask_train].cpu().flatten(), label[mask_train].cpu().flatten()
)
optimizer.zero_grad()
loss.backward()
optimizer.step()
optim_scheduler.step()
# Evaluate on training subset for early stopping
model.eval()
train_auc, _ = get_metrics_auc(
label[mask_train].cpu().detach().numpy(),
pred[mask_train].cpu().detach().numpy(),
)
early_stop = stopper.step(loss.item(), train_auc, model)
# Every 50 epochs, print test metrics (on full graph, no sampling)
if epoch % 50 == 0:
test_auc, test_aupr = get_metrics_auc(
label[mask_test].cpu().detach().numpy(),
pred[mask_test].cpu().detach().numpy(),
)
time_end = time.time()
time_cost = (time_end - time_start) / 3600
time_remaining = (time_cost / (epoch / 50)) * (args.epoch - epoch)
print(
f"Epoch {epoch} Loss: {loss.item():.3f}; Train AUC: {train_auc:.3f}; "
f"Test AUC: {test_auc:.3f}; Test AUPR: {test_aupr:.3f}; Est. Time Remaining: {time_remaining:.2f}h"
)
print("-" * 50)
# if early_stop:
# break
# Load best checkpoint and produce final predictions on full graph
stopper.load_checkpoint(model)
model.eval()
# save the model
th.save(model.state_dict(), os.path.join(args.saved_path, f"model_{fold}.pth"))
final_pred = th.sigmoid(model(g, feature)).cpu().detach().numpy()
# print(final_pred.shape)
# print(pred_result.shape)
pred_result[test_pos_indices[0], test_pos_indices[1]] = final_pred[test_pos_indices[0], test_pos_indices[1]]
pred_result[test_neg_indices[0], test_neg_indices[1]] = final_pred[test_neg_indices[0], test_neg_indices[1]]
# save pred_result to numpy array
# np.save(os.path.join(args.saved_path, f"pred_result_{fold}.npy"), final_pred)
# save indices to numpy array
# np.save(os.path.join(args.saved_path, f"test_pos_indices_{fold}.npy"), test_pos_indices)
# np.save(os.path.join(args.saved_path, f"test_neg_indices_{fold}.npy"), test_neg_indices)
fold += 1
# Overall evaluation
# overall_auc, overall_aupr, overall_acc, overall_f1, overall_pre, overall_rec = get_metrics(
# label.cpu().detach().numpy().flatten(), pred_result.flatten()
# )
# print(
# f"Overall: AUC {overall_auc:.3f}; AUPR {overall_aupr:.3f}; Acc {overall_acc:.3f}; "
# f"F1 {overall_f1:.3f}; Precision {overall_pre:.3f}; Recall {overall_rec:.3f};"
# )
# Save predictions
result_path = os.path.join(args.saved_path, "result.csv")
pd.DataFrame(pred_result).to_csv(result_path, index=False, header=False)
# Plot curves
# plot_result_auc(args, data[:, -1].flatten(), pred_result.flatten(), overall_auc)
# plot_result_aupr(args, data[:, -1].flatten(), pred_result.flatten(), overall_aupr)
if __name__ == "__main__":
run_training()