Files
NexQuant/rdagent/scenarios/data_mining/experiment/model_template/train.py
T

117 lines
3.0 KiB
Python
Raw Normal View History

import os
import random
2024-07-24 16:56:27 +08:00
from pathlib import Path
import numpy as np
2024-07-24 16:56:27 +08:00
import pandas as pd
import sparse
2024-07-24 16:56:27 +08:00
import torch
import torch.nn as nn
import torch.nn.functional as F
2024-07-24 16:56:27 +08:00
from model import model_cls
from sklearn.metrics import accuracy_score, roc_auc_score
from torch.utils.data import DataLoader, Dataset
from torchvision import datasets, transforms
2024-07-24 16:56:27 +08:00
# Set device for training
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# device = torch.device("cpu")
2024-07-24 16:56:27 +08:00
class MyDataset(Dataset):
def __init__(self, x, label, device):
self.x1 = x
self.label = label
self.device = device
def __len__(self):
return len(self.label)
def __getitem__(self, idx):
if torch.is_tensor(idx):
idx = idx.tolist()
return torch.FloatTensor(self.x1[idx]).to(self.device), torch.tensor(self.label[idx], dtype=torch.float).to(
self.device
)
2024-07-24 16:56:27 +08:00
def collate_fn(batch):
x, label = [], []
for data in batch:
x.append(data[0])
label.append(data[1])
return torch.stack(x, 0), torch.stack(label, 0)
datapath = "/root/.data"
2024-07-24 16:56:27 +08:00
# datapath = '/home/v-suhancui/RD-Agent/physionet.org/files/mimic-eicu-fiddle-feature/1.0.0/FIDDLE_mimic3'
X = sparse.load_npz(datapath + "/features/ARF_12h/X.npz").todense()
df_pop = pd.read_csv(datapath + "/population/ARF_12h.csv")["ARF_LABEL"]
2024-07-24 16:56:27 +08:00
X = X.transpose(0, 2, 1)
indices = [i for i in range(len(df_pop))]
random.shuffle(indices)
split_point = int(0.7 * len(df_pop))
X_train, y_train = X[indices[:split_point]], np.array(df_pop[indices[:split_point]])
X_test, y_test = X[indices[split_point:]], np.array(df_pop[indices[split_point:]])
train_dataloader = DataLoader(
MyDataset(X_train, y_train, device), collate_fn=collate_fn, shuffle=True, drop_last=True, batch_size=64
)
test_dataloader = DataLoader(
MyDataset(X_test, y_test, device), collate_fn=collate_fn, shuffle=False, drop_last=False, batch_size=64
)
2024-07-24 16:56:27 +08:00
num_features = 4816
num_timesteps = 12
# Define the optimizer and loss function
model = model_cls(num_features=num_features, num_timesteps=num_timesteps).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.0001)
criterion = nn.CrossEntropyLoss()
2024-08-01 15:00:40 +08:00
2024-07-24 16:56:27 +08:00
# Train the model
2024-08-01 15:00:40 +08:00
def eval_auc(model):
y_pred = []
for data in test_dataloader:
x, y = data
out = model(x)
y_pred.append(out.cpu().detach().numpy())
return roc_auc_score(y_test, np.concatenate(y_pred))
best = 0.0
best_model = None
2024-07-24 16:56:27 +08:00
2024-08-01 15:00:40 +08:00
for i in range(15):
2024-07-24 16:56:27 +08:00
for data in train_dataloader:
x, y = data
out = model(x)
optimizer.zero_grad()
loss = criterion(out.squeeze(), y)
loss.backward()
optimizer.step()
2024-08-01 15:00:40 +08:00
roc = eval_auc(model)
if roc > best:
best = roc
best_model = model
2024-07-24 16:56:27 +08:00
y_pred = []
for data in test_dataloader:
x, y = data
2024-08-01 15:00:40 +08:00
out = best_model(x)
2024-07-24 16:56:27 +08:00
y_pred.append(out.cpu().detach().numpy())
acc = roc_auc_score(y_test, np.concatenate(y_pred))
print(acc)
2024-07-30 11:38:02 +08:00
2024-07-30 12:14:48 +08:00
res = pd.Series(data=[acc], index=["AUROC"])
2024-07-30 11:38:02 +08:00
res.to_csv("./submission.csv")