| import torch |
| import torch.nn as nn |
| import torch.optim as optim |
| from torch.utils.data import DataLoader |
| from sklearn.metrics import roc_auc_score, f1_score, classification_report |
| import joblib |
| import os |
|
|
| |
| from model import HybridTabTransformer |
| from dataset import HeartDiseaseDataset |
|
|
| def train_model(): |
| |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| metadata = joblib.load('assets/model_metadata.joblib') |
| batch_size = 32 |
| epochs = 20 |
| lr = 0.001 |
|
|
| |
| train_ds = HeartDiseaseDataset('data/processed/train.csv') |
| test_ds = HeartDiseaseDataset('data/processed/test.csv') |
| |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True) |
| test_loader = DataLoader(test_ds, batch_size=batch_size) |
|
|
| |
| model = HybridTabTransformer( |
| cat_dims=metadata['cat_dims'], |
| num_continuous=len(metadata['num_cols']) |
| ).to(device) |
|
|
| criterion = nn.BCELoss() |
| optimizer = optim.Adam(model.parameters(), lr=lr) |
|
|
| print(f"Starting training on {device}...") |
|
|
| |
| for epoch in range(epochs): |
| model.train() |
| total_loss = 0 |
| for x_cat, x_num, y in train_loader: |
| x_cat, x_num, y = x_cat.to(device), x_num.to(device), y.to(device) |
| |
| optimizer.zero_grad() |
| outputs = model(x_cat, x_num) |
| loss = criterion(outputs, y) |
| loss.backward() |
| optimizer.step() |
| total_loss += loss.item() |
| |
| if (epoch + 1) % 5 == 0: |
| print(f"Epoch [{epoch+1}/{epochs}], Loss: {total_loss/len(train_loader):.4f}") |
|
|
| |
| model.eval() |
| all_preds = [] |
| all_targets = [] |
| |
| with torch.no_grad(): |
| for x_cat, x_num, y in test_loader: |
| x_cat, x_num, y = x_cat.to(device), x_num.to(device), y.to(device) |
| outputs = model(x_cat, x_num) |
| all_preds.extend(outputs.cpu().numpy()) |
| all_targets.extend(y.cpu().numpy()) |
|
|
| |
| binary_preds = [1 if p >= 0.5 else 0 for p in all_preds] |
| |
| print("\n--- Final Model Evaluation ---") |
| print(f"AUROC Score: {roc_auc_score(all_targets, all_preds):.4f}") |
| print(f"F1 Score: {f1_score(all_targets, binary_preds):.4f}") |
| print("\nClassification Report:") |
| print(classification_report(all_targets, binary_preds)) |
|
|
| |
| torch.save(model.state_dict(), 'assets/model.pth') |
| print("Model saved to assets/model.pth") |
|
|
| if __name__ == "__main__": |
| train_model() |