File size: 3,555 Bytes
43a9675
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
import os
import argparse
from datetime import datetime

import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from tqdm import tqdm
from sklearn.metrics import roc_auc_score, f1_score

from dataset import PTBXLDataset
from models.hmt_ecgnet import HMT_ECGNet
from utils_seed import set_seed
from config import BATCH_SIZE, LR, WEIGHT_DECAY, N_EPOCHS, N_LEADS


# ---------------- ARGUMENTS ----------------
parser = argparse.ArgumentParser()
parser.add_argument("--task", choices=["mi_vs_norm", "norm_vs_abnormal"], required=True)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--epochs", type=int, default=N_EPOCHS)
args = parser.parse_args()


def get_loader(ds, shuffle):
    use_cuda = torch.cuda.is_available()
    return DataLoader(
        ds,
        batch_size=BATCH_SIZE,
        shuffle=shuffle,
        num_workers=0,
        pin_memory=use_cuda
    )


def main():
    set_seed(args.seed)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    run_name = f"hmt_binary_{args.task}_seed{args.seed}_{datetime.now():%Y%m%d-%H%M%S}"
    log_dir = os.path.join("logs", run_name)
    os.makedirs(log_dir, exist_ok=True)

    # ---------------- DATA ----------------
    train_ds = PTBXLDataset(split="train", task="binary", binary_task=args.task)
    val_ds   = PTBXLDataset(split="val",   task="binary", binary_task=args.task)

    train_loader = get_loader(train_ds, shuffle=True)
    val_loader   = get_loader(val_ds, shuffle=False)

    # ---------------- MODEL ----------------
    model = HMT_ECGNet(num_classes=1, num_leads=N_LEADS).to(device)
    print(f"Model parameters: {sum(p.numel() for p in model.parameters()):,}")

    criterion = nn.BCEWithLogitsLoss()
    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)

    best_auc = 0.0

    for epoch in range(1, args.epochs + 1):

        # -------- TRAIN --------
        model.train()
        train_loss = 0.0

        for x, y in tqdm(train_loader, desc=f"Epoch {epoch:03d} [Train]"):
            x, y = x.to(device), y.to(device).float()

            optimizer.zero_grad()
            logits = model(x).squeeze(1)
            loss = criterion(logits, y)
            loss.backward()
            optimizer.step()

            train_loss += loss.item()

        train_loss /= len(train_loader)

        # -------- VALIDATE --------
        model.eval()
        all_logits, all_targets = [], []

        with torch.no_grad():
            for x, y in val_loader:
                x = x.to(device)
                logits = model(x).squeeze(1)
                all_logits.append(logits.cpu())
                all_targets.append(y)

        logits = torch.cat(all_logits)
        targets = torch.cat(all_targets)

        probs = torch.sigmoid(logits).numpy()
        val_auc = roc_auc_score(targets.numpy(), probs)
        val_f1  = f1_score(targets.numpy(), probs >= 0.5)

        print(
            f"Epoch {epoch:03d} | loss={train_loss:.4f} | "
            f"AUROC={val_auc:.4f} | F1={val_f1:.4f}"
        )

        if val_auc > best_auc:
            best_auc = val_auc
            torch.save(
                {"epoch": epoch, "model_state_dict": model.state_dict()},
                os.path.join(log_dir, f"hmt_binary_{args.task}_best.pth"),
            )
            print("  -> Best model saved")

    print("Training complete.")


if __name__ == "__main__":
    main()