summaryrefslogtreecommitdiff
path: root/train.py
blob: 545b299782f06342231736b34a452dcc50802e03 (plain)
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
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
import copy
import time

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


def train_model(
    model: nn.Module,
    train_loader: DataLoader,
    val_loader: DataLoader,
    criterion: nn.Module,
    optimizer: torch.optim.Optimizer,
    scheduler,
    device: torch.device,
    num_epochs: int = 30,
    patience: int = 5,
    model_name: str = "",
) -> dict:
    """Train model with mixed precision, early stopping on val weighted-F1.

    Returns a dict with:
      best_model_state, train_losses, val_losses, val_f1s,
      train_time (seconds), epochs_trained
    """
    scaler = torch.amp.GradScaler('cuda') if device.type == 'cuda' else None

    best_val_f1 = -1.0
    best_state = None
    patience_counter = 0

    train_losses, val_losses, val_f1s = [], [], []
    start_time = time.time()

    for epoch in range(1, num_epochs + 1):
        # ── Training phase ──────────────────────────────────────────────────
        model.train()
        running_loss = 0.0
        n_train = 0

        for images, labels in train_loader:
            images = images.to(device, non_blocking=True)
            labels = labels.to(device, non_blocking=True)

            optimizer.zero_grad()

            if scaler is not None:
                with torch.amp.autocast('cuda'):
                    outputs = model(images)
                    loss = criterion(outputs, labels)
                scaler.scale(loss).backward()
                scaler.unscale_(optimizer)
                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
                scaler.step(optimizer)
                scaler.update()
            else:
                outputs = model(images)
                loss = criterion(outputs, labels)
                loss.backward()
                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
                optimizer.step()

            running_loss += loss.item() * images.size(0)
            n_train += images.size(0)

        train_loss = running_loss / n_train
        train_losses.append(train_loss)

        # ── Validation phase ─────────────────────────────────────────────────
        model.eval()
        val_running_loss = 0.0
        n_val = 0
        all_preds, all_labels = [], []

        with torch.no_grad():
            for images, labels in val_loader:
                images = images.to(device, non_blocking=True)
                labels = labels.to(device, non_blocking=True)

                if scaler is not None:
                    with torch.amp.autocast('cuda'):
                        outputs = model(images)
                        loss = criterion(outputs, labels)
                else:
                    outputs = model(images)
                    loss = criterion(outputs, labels)

                val_running_loss += loss.item() * images.size(0)
                n_val += images.size(0)

                preds = outputs.argmax(dim=1)
                all_preds.extend(preds.cpu().numpy())
                all_labels.extend(labels.cpu().numpy())

        val_loss = val_running_loss / n_val
        val_losses.append(val_loss)

        val_f1 = f1_score(all_labels, all_preds, average='weighted', zero_division=0)
        val_f1s.append(val_f1)

        current_lr = scheduler.get_last_lr()[0] if hasattr(scheduler, 'get_last_lr') else optimizer.param_groups[0]['lr']
        print(
            f"[{model_name}] Epoch {epoch:3d}/{num_epochs} | "
            f"train loss: {train_loss:.4f} | val loss: {val_loss:.4f} | "
            f"val F1: {val_f1:.4f} | lr: {current_lr:.2e}"
        )

        scheduler.step()

        # ── Early stopping ────────────────────────────────────────────────────
        if val_f1 > best_val_f1:
            best_val_f1 = val_f1
            best_state = copy.deepcopy(model.state_dict())
            patience_counter = 0
        else:
            patience_counter += 1
            if patience_counter >= patience:
                print(f"  Early stopping at epoch {epoch} (no improvement for {patience} epochs).")
                break

    train_time = time.time() - start_time
    epochs_trained = len(train_losses)

    print(f"  Training complete: {epochs_trained} epochs, {train_time:.1f}s, best val F1: {best_val_f1:.4f}")

    return {
        'best_model_state': best_state,
        'train_losses': train_losses,
        'val_losses': val_losses,
        'val_f1s': val_f1s,
        'train_time': train_time,
        'epochs_trained': epochs_trained,
        'best_val_f1': best_val_f1,
    }