summaryrefslogtreecommitdiff
path: root/main.py
diff options
context:
space:
mode:
authorVineet Kumar <git@vineetk.net>2026-04-20 10:03:17 -0400
committerVineet Kumar <git@vineetk.net>2026-04-20 10:03:17 -0400
commit9dd8b3007e79f0634f868d405f5e7c48917305fd (patch)
treeb2dd74419a78b28190664f482b30263879742d7a /main.py
parent8f571ee51854540d186aafc01d25908bba5ea09e (diff)
improve training for class-imbalanced DR classification
- WeightedRandomSampler for balanced class sampling during training - FocalLoss (gamma=2) for hard-example mining, no alpha (sampler handles balance) - Stronger augmentation: RandomResizedCrop, ColorJitter, GaussianBlur, 30deg rotation - Discriminative LR: backbone at 0.1x, head at 1x - 3-epoch linear warmup + cosine decay scheduler - Early stopping on macro F1 instead of weighted F1 - Updated defaults: batch_size=64, epochs=50, patience=10 - Comparison chart now shows both macro and weighted F1 - DataParallel checkpoint handling (strip module. prefix) - Enable cudnn.benchmark for training speed
Diffstat (limited to 'main.py')
-rw-r--r--main.py60
1 files changed, 43 insertions, 17 deletions
diff --git a/main.py b/main.py
index 6d8b2d9..fb87c33 100644
--- a/main.py
+++ b/main.py
@@ -12,10 +12,10 @@ import os
12 12
13import torch 13import torch
14import torch.nn as nn 14import torch.nn as nn
15from torch.optim.lr_scheduler import CosineAnnealingLR 15from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR
16import kagglehub 16import kagglehub
17 17
18from utils import set_seed, compute_class_weights, get_device, CLASS_NAMES 18from utils import set_seed, get_device, FocalLoss
19from dataset import split_dataset, build_dataloaders 19from dataset import split_dataset, build_dataloaders
20from models import create_model 20from models import create_model
21from train import train_model 21from train import train_model
@@ -35,15 +35,17 @@ def parse_args():
35 '--clahe', nargs='+', type=int, default=[0, 1], 35 '--clahe', nargs='+', type=int, default=[0, 1],
36 help='CLAHE settings to run (0=off, 1=on)' 36 help='CLAHE settings to run (0=off, 1=on)'
37 ) 37 )
38 parser.add_argument('--epochs', type=int, default=30, help='Max training epochs') 38 parser.add_argument('--epochs', type=int, default=50, help='Max training epochs')
39 parser.add_argument('--patience', type=int, default=5, help='Early stopping patience') 39 parser.add_argument('--patience', type=int, default=10, help='Early stopping patience')
40 parser.add_argument('--batch-size', type=int, default=32, help='Training batch size') 40 parser.add_argument('--batch-size', type=int, default=64, help='Training batch size')
41 parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate') 41 parser.add_argument('--lr', type=float, default=1e-4, help='Learning rate')
42 parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay') 42 parser.add_argument('--weight-decay', type=float, default=1e-4, help='AdamW weight decay')
43 parser.add_argument('--seed', type=int, default=42, help='Random seed') 43 parser.add_argument('--seed', type=int, default=42, help='Random seed')
44 parser.add_argument('--data-root', type=str, default=None, 44 parser.add_argument('--data-root', type=str, default=None,
45 help='Path to dataset root (auto-downloaded if not set)') 45 help='Path to dataset root (auto-downloaded if not set)')
46 parser.add_argument('--num-workers', type=int, default=4, help='DataLoader workers') 46 parser.add_argument('--num-workers', type=int, default=8, help='DataLoader workers')
47 parser.add_argument('--focal-gamma', type=float, default=2.0, help='Focal loss gamma (0=standard CE)')
48 parser.add_argument('--warmup-epochs', type=int, default=3, help='LR warmup epochs')
47 return parser.parse_args() 49 return parser.parse_args()
48 50
49 51
@@ -64,10 +66,6 @@ def main():
64 train_paths, train_labels, val_paths, val_labels, test_paths, test_labels = \ 66 train_paths, train_labels, val_paths, val_labels, test_paths, test_labels = \
65 split_dataset(data_root, seed=args.seed) 67 split_dataset(data_root, seed=args.seed)
66 68
67 # Class weights computed once from training labels
68 class_weights = compute_class_weights(train_labels).to(device)
69 print("Class weights:", [f"{w:.3f}" for w in class_weights.cpu().numpy()])
70
71 # ── Experiments ────────────────────────────────────────────────────────── 69 # ── Experiments ──────────────────────────────────────────────────────────
72 experiments = [ 70 experiments = [
73 (model_name, bool(use_clahe)) 71 (model_name, bool(use_clahe))
@@ -97,11 +95,35 @@ def main():
97 ) 95 )
98 96
99 model = create_model(model_name, num_classes=5, pretrained=True).to(device) 97 model = create_model(model_name, num_classes=5, pretrained=True).to(device)
100 criterion = nn.CrossEntropyLoss(weight=class_weights) 98
99 # Discriminative LR: backbone gets 10x lower LR than the new head
100 if model_name == 'resnet50':
101 head_params = list(model.fc.parameters())
102 elif model_name == 'efficientnet_b0':
103 head_params = list(model.classifier.parameters())
104 elif model_name == 'vit_b_16':
105 head_params = list(model.heads.parameters())
106 else:
107 head_params = []
108 head_ids = {id(p) for p in head_params}
109 backbone_params = [p for p in model.parameters() if id(p) not in head_ids]
110
111 if torch.cuda.device_count() > 1:
112 print(f" Using {torch.cuda.device_count()} GPUs via DataParallel")
113 model = torch.nn.DataParallel(model)
114
115 criterion = FocalLoss(alpha=None, gamma=args.focal_gamma)
101 optimizer = torch.optim.AdamW( 116 optimizer = torch.optim.AdamW(
102 model.parameters(), lr=args.lr, weight_decay=args.weight_decay 117 [
118 {'params': backbone_params, 'lr': args.lr * 0.1},
119 {'params': head_params, 'lr': args.lr},
120 ],
121 weight_decay=args.weight_decay,
103 ) 122 )
104 scheduler = CosineAnnealingLR(optimizer, T_max=args.epochs) 123 warmup_epochs = min(args.warmup_epochs, args.epochs - 1)
124 warmup_scheduler = LinearLR(optimizer, start_factor=0.1, total_iters=warmup_epochs)
125 cosine_scheduler = CosineAnnealingLR(optimizer, T_max=max(args.epochs - warmup_epochs, 1))
126 scheduler = SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[warmup_epochs])
105 127
106 # Train 128 # Train
107 train_results = train_model( 129 train_results = train_model(
@@ -113,16 +135,20 @@ def main():
113 model_name=exp_key, 135 model_name=exp_key,
114 ) 136 )
115 137
116 # Save best checkpoint 138 # Save best checkpoint (strip DataParallel 'module.' prefix for portability)
117 ckpt_path = os.path.join(exp_dir, 'best_model.pth') 139 ckpt_path = os.path.join(exp_dir, 'best_model.pth')
118 torch.save(train_results['best_model_state'], ckpt_path) 140 best_state = train_results['best_model_state']
141 if isinstance(model, torch.nn.DataParallel):
142 best_state = {k[7:]: v for k, v in best_state.items()}
143 torch.save(best_state, ckpt_path)
119 print(f" Checkpoint saved: {ckpt_path}") 144 print(f" Checkpoint saved: {ckpt_path}")
120 145
121 # Load best weights for evaluation 146 # Load best weights for evaluation
122 model.load_state_dict(train_results['best_model_state']) 147 core_model = model.module if isinstance(model, torch.nn.DataParallel) else model
148 core_model.load_state_dict(best_state)
123 149
124 # Evaluate on test set 150 # Evaluate on test set
125 test_metrics = evaluate_model(model, test_loader, device) 151 test_metrics = evaluate_model(core_model, test_loader, device)
126 152
127 print(f"\n Test results — Weighted F1: {test_metrics['weighted_f1']:.4f} | " 153 print(f"\n Test results — Weighted F1: {test_metrics['weighted_f1']:.4f} | "
128 f"Macro F1: {test_metrics['macro_f1']:.4f} | " 154 f"Macro F1: {test_metrics['macro_f1']:.4f} | "