summaryrefslogtreecommitdiff
path: root/main.py
blob: 0d29d634cd3fe7238e3a6c4a6157102db8832bfd (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
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
"""
EEL4759 Final Project — Diabetic Retinopathy Classification
Compares ResNet-50, EfficientNet-B0, and ViT-B/16 with and without CLAHE preprocessing.

Usage:
    python main.py                          # run all 6 experiments (default)
    python main.py --models resnet50 --clahe 0 --epochs 2   # quick sanity check
"""

import argparse
import os

import torch
import torch.nn as nn
from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR
import kagglehub

from utils import set_seed, get_device, FocalLoss
from dataset import split_dataset, build_dataloaders
from models import create_model
from train import train_model
from evaluate import evaluate_model, plot_confusion_matrix, plot_training_curves, generate_summary

RESULTS_DIR = 'results'


def parse_args():
    parser = argparse.ArgumentParser(description='Diabetic Retinopathy Classification')
    parser.add_argument(
        '--models', nargs='+',
        default=['resnet50', 'efficientnet_b0', 'vit_b_16'],
        help='Models to run'
    )
    parser.add_argument(
        '--clahe', nargs='+', type=int, default=[0, 1],
        help='CLAHE settings to run (0=off, 1=on)'
    )
    parser.add_argument('--epochs', type=int, default=50, help='Max training epochs')
    parser.add_argument('--patience', type=int, default=10, help='Early stopping patience')
    parser.add_argument('--batch-size', type=int, default=256, help='Training batch size')
    parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate')
    parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay')
    parser.add_argument('--seed', type=int, default=42, help='Random seed')
    parser.add_argument('--data-root', type=str, default=None,
                        help='Path to dataset root (auto-downloaded if not set)')
    parser.add_argument('--num-workers', type=int, default=8, help='DataLoader workers')
    parser.add_argument('--focal-gamma', type=float, default=2.0, help='Focal loss gamma (0=standard CE)')
    parser.add_argument('--warmup-epochs', type=int, default=3, help='LR warmup epochs')
    parser.add_argument('--img-size', type=int, default=384, help='Input image size')
    return parser.parse_args()


def main():
    args = parse_args()
    set_seed(args.seed)
    device = get_device()
    print(f"Using device: {device}")

    # ── Dataset ──────────────────────────────────────────────────────────────
    if args.data_root is None:
        print("Downloading/locating dataset via kagglehub...")
        data_root = kagglehub.dataset_download("amanneo/diabetic-retinopathy-resized-arranged")
    else:
        data_root = args.data_root
    print(f"Dataset root: {data_root}")

    train_paths, train_labels, val_paths, val_labels, test_paths, test_labels = \
        split_dataset(data_root, seed=args.seed)

    # ── Experiments ──────────────────────────────────────────────────────────
    experiments = [
        (model_name, bool(use_clahe))
        for model_name in args.models
        for use_clahe in args.clahe
    ]

    all_results = {}
    os.makedirs(RESULTS_DIR, exist_ok=True)

    for model_name, use_clahe in experiments:
        exp_key = f"{model_name}_clahe={use_clahe}"
        exp_dir = os.path.join(RESULTS_DIR, exp_key)
        os.makedirs(exp_dir, exist_ok=True)

        print(f"\n{'='*70}")
        print(f"Experiment: {model_name}  |  CLAHE: {use_clahe}")
        print(f"{'='*70}")

        train_loader, val_loader, test_loader = build_dataloaders(
            train_paths, train_labels,
            val_paths, val_labels,
            test_paths, test_labels,
            use_clahe=use_clahe,
            batch_size=args.batch_size,
            num_workers=args.num_workers,
            img_size=args.img_size,
        )

        model = create_model(model_name, num_classes=5, pretrained=True).to(device)

        # Discriminative LR: backbone gets 10x lower LR than the new head
        if model_name == 'resnet50':
            head_params = list(model.fc.parameters())
        elif model_name == 'efficientnet_b0':
            head_params = list(model.classifier.parameters())
        elif model_name == 'vit_b_16':
            head_params = list(model.heads.parameters())
        else:
            head_params = []
        head_ids = {id(p) for p in head_params}
        backbone_params = [p for p in model.parameters() if id(p) not in head_ids]

        if torch.cuda.device_count() > 1:
            print(f"  Using {torch.cuda.device_count()} GPUs via DataParallel")
            model = torch.nn.DataParallel(model)

        criterion = FocalLoss(alpha=None, gamma=args.focal_gamma)
        optimizer = torch.optim.AdamW(
            [
                {'params': backbone_params, 'lr': args.lr * 0.1},
                {'params': head_params, 'lr': args.lr},
            ],
            weight_decay=args.weight_decay,
        )
        warmup_epochs = min(args.warmup_epochs, args.epochs - 1)
        warmup_scheduler = LinearLR(optimizer, start_factor=0.1, total_iters=warmup_epochs)
        cosine_scheduler = CosineAnnealingLR(optimizer, T_max=max(args.epochs - warmup_epochs, 1))
        scheduler = SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[warmup_epochs])

        # Train
        train_results = train_model(
            model, train_loader, val_loader,
            criterion, optimizer, scheduler,
            device,
            num_epochs=args.epochs,
            patience=args.patience,
            model_name=exp_key,
        )

        # Save best checkpoint (strip DataParallel 'module.' prefix for portability)
        ckpt_path = os.path.join(exp_dir, 'best_model.pth')
        best_state = train_results['best_model_state']
        if isinstance(model, torch.nn.DataParallel):
            best_state = {k[7:]: v for k, v in best_state.items()}
        torch.save(best_state, ckpt_path)
        print(f"  Checkpoint saved: {ckpt_path}")

        # Load best weights for evaluation
        core_model = model.module if isinstance(model, torch.nn.DataParallel) else model
        core_model.load_state_dict(best_state)

        # Evaluate on test set
        test_metrics = evaluate_model(core_model, test_loader, device)

        print(f"\n  Test results — Weighted F1: {test_metrics['weighted_f1']:.4f} | "
              f"Macro F1: {test_metrics['macro_f1']:.4f} | "
              f"Accuracy: {test_metrics['accuracy']:.4f} | "
              f"Kappa: {test_metrics['cohen_kappa']:.4f}")
        print(test_metrics['classification_report'])

        # Save per-class report
        report_path = os.path.join(exp_dir, 'classification_report.txt')
        with open(report_path, 'w') as f:
            f.write(f"Experiment: {exp_key}\n\n")
            f.write(test_metrics['classification_report'])

        # Plots
        plot_confusion_matrix(
            test_metrics['confusion_matrix'],
            save_path=os.path.join(exp_dir, 'confusion_matrix.png'),
        )
        plot_training_curves(
            train_results['train_losses'],
            train_results['val_losses'],
            train_results['val_f1s'],
            save_path=os.path.join(exp_dir, 'training_curves.png'),
        )

        all_results[exp_key] = {**train_results, **test_metrics}
        # Remove bulky state dict from summary dict (already saved to disk)
        all_results[exp_key].pop('best_model_state', None)

    # ── Summary ───────────────────────────────────────────────────────────────
    if len(all_results) > 1:
        generate_summary(all_results, save_dir=RESULTS_DIR)


if __name__ == '__main__':
    main()