diff options
Diffstat (limited to 'main.py')
| -rw-r--r-- | main.py | 60 |
1 files changed, 43 insertions, 17 deletions
| @@ -12,10 +12,10 @@ import os | |||
| 12 | 12 | ||
| 13 | import torch | 13 | import torch |
| 14 | import torch.nn as nn | 14 | import torch.nn as nn |
| 15 | from torch.optim.lr_scheduler import CosineAnnealingLR | 15 | from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR |
| 16 | import kagglehub | 16 | import kagglehub |
| 17 | 17 | ||
| 18 | from utils import set_seed, compute_class_weights, get_device, CLASS_NAMES | 18 | from utils import set_seed, get_device, FocalLoss |
| 19 | from dataset import split_dataset, build_dataloaders | 19 | from dataset import split_dataset, build_dataloaders |
| 20 | from models import create_model | 20 | from models import create_model |
| 21 | from train import train_model | 21 | from 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} | " |
