""" train.py – Training loop for Chest X-Ray classification. Usage: # Train MobileNetV2 (default 15 epochs) python src/train.py --model mobilenet_v2 # Train SimpleCNN baseline for 20 epochs python src/train.py --model simple_cnn --epochs 20 # Quick smoke test (2 epochs) python src/train.py --model mobilenet_v2 --epochs 2 """ import argparse import logging import random import sys import time from pathlib import Path import numpy as np import torch import torch.nn as nn import yaml from torch.optim import Adam from torch.optim.lr_scheduler import ReduceLROnPlateau from tqdm import tqdm # Allow running from project root or from src/ sys.path.insert(0, str(Path(__file__).parent)) from dataset import build_dataloaders, LABEL_NAMES, LABEL_KEY from model import build_model, MobileNetV2Classifier # ─── Logging ────────────────────────────────────────────────────────────────── logging.basicConfig( level=logging.INFO, format="%(asctime)s | %(levelname)-8s | %(name)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S", ) logger = logging.getLogger(__name__) # ─── Reproducibility ────────────────────────────────────────────────────────── def set_seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) # ─── Per-batch helpers ──────────────────────────────────────────────────────── def _accuracy(logits: torch.Tensor, labels: torch.Tensor) -> float: return (logits.argmax(dim=1) == labels).float().mean().item() def train_one_epoch(model, loader, optimizer, criterion, device) -> dict: model.train() total_loss = total_acc = 0.0 for images, labels in tqdm(loader, desc=" train", leave=False): images, labels = images.to(device), labels.to(device) optimizer.zero_grad() logits = model(images) loss = criterion(logits, labels) loss.backward() optimizer.step() total_loss += loss.item() total_acc += _accuracy(logits, labels) n = len(loader) return {"loss": total_loss / n, "acc": total_acc / n} @torch.no_grad() def eval_one_epoch(model, loader, criterion, device) -> dict: model.eval() total_loss = total_acc = 0.0 for images, labels in tqdm(loader, desc=" val ", leave=False): images, labels = images.to(device), labels.to(device) logits = model(images) loss = criterion(logits, labels) total_loss += loss.item() total_acc += _accuracy(logits, labels) n = len(loader) return {"loss": total_loss / n, "acc": total_acc / n} # ─── Main training function ─────────────────────────────────────────────────── def train(cfg: dict, model_name_override: str = None, epochs_override: int = None) -> dict: """ Full training loop. Args: cfg: Config dict from config.yaml. model_name_override: Override cfg['model']['name'] if provided. epochs_override: Override cfg['training']['epochs'] if provided. Returns: Training history dict (train_loss, train_acc, val_loss, val_acc per epoch). """ if model_name_override: cfg["model"]["name"] = model_name_override if epochs_override: cfg["training"]["epochs"] = epochs_override t_cfg = cfg["training"] set_seed(t_cfg["seed"]) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") logger.info(f"Device: {device}") # ── Data ───────────────────────────────────────────────────────────────── loaders = build_dataloaders(cfg) # ── Model ──────────────────────────────────────────────────────────────── model = build_model(cfg).to(device) # ── Loss: class-weighted CE to address imbalance ───────────────────────── raw_labels = [loaders["train"].dataset.data[i][LABEL_KEY] for i in range(len(loaders["train"].dataset))] class_counts = np.bincount(raw_labels) class_weights = torch.tensor(1.0 / class_counts.astype(float), dtype=torch.float).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights) logger.info( "Class weights: " + ", ".join(f"{LABEL_NAMES[i]}={w:.4f}" for i, w in enumerate(class_weights.cpu())) ) # ── Optimiser & scheduler ───────────────────────────────────────────────── optimizer = Adam( filter(lambda p: p.requires_grad, model.parameters()), lr=t_cfg["learning_rate"], weight_decay=t_cfg["weight_decay"], ) scheduler = ReduceLROnPlateau( optimizer, mode="max", patience=t_cfg["patience"], factor=0.3 ) # ── Checkpointing ───────────────────────────────────────────────────────── ckpt_dir = Path(t_cfg["checkpoint_dir"]) ckpt_dir.mkdir(parents=True, exist_ok=True) model_name = cfg["model"]["name"] best_ckpt = ckpt_dir / f"best_{model_name}.pth" best_val_acc = 0.0 unfreeze_after = cfg["model"].get("unfreeze_after_epoch") history: dict = {k: [] for k in ["train_loss", "train_acc", "val_loss", "val_acc"]} logger.info(f"Starting training: model={model_name}, epochs={t_cfg['epochs']}") logger.info("─" * 70) for epoch in range(1, t_cfg["epochs"] + 1): t0 = time.time() # Unfreeze backbone after N epochs (gradual fine-tuning) if ( unfreeze_after and epoch == unfreeze_after and isinstance(model, MobileNetV2Classifier) ): model.unfreeze_backbone() # Rebuild optimiser with lower LR for backbone parameters fine_tune_lr = t_cfg["learning_rate"] * 0.1 optimizer = Adam( model.parameters(), lr=fine_tune_lr, weight_decay=t_cfg["weight_decay"] ) scheduler = ReduceLROnPlateau( optimizer, mode="max", patience=t_cfg["patience"], factor=0.3 ) logger.info(f"Epoch {epoch}: backbone unfrozen, LR={fine_tune_lr:.2e}") train_m = train_one_epoch(model, loaders["train"], optimizer, criterion, device) val_m = eval_one_epoch(model, loaders["val"], criterion, device) scheduler.step(val_m["acc"]) elapsed = time.time() - t0 logger.info( f"Epoch {epoch:03d}/{t_cfg['epochs']:03d} " f"train loss={train_m['loss']:.4f} acc={train_m['acc']:.4f} " f"val loss={val_m['loss']:.4f} acc={val_m['acc']:.4f} " f"({elapsed:.1f}s)" ) history["train_loss"].append(train_m["loss"]) history["train_acc"].append(train_m["acc"]) history["val_loss"].append(val_m["loss"]) history["val_acc"].append(val_m["acc"]) if val_m["acc"] > best_val_acc: best_val_acc = val_m["acc"] torch.save( { "epoch": epoch, "model_state": model.state_dict(), "val_acc": best_val_acc, "cfg": cfg, }, best_ckpt, ) logger.info(f" ✓ Best model saved (val_acc={best_val_acc:.4f})") logger.info("─" * 70) logger.info(f"Training done. Best val_acc={best_val_acc:.4f} → {best_ckpt}") return history # ─── CLI entry ──────────────────────────────────────────────────────────────── if __name__ == "__main__": parser = argparse.ArgumentParser(description="Train Chest X-Ray classifier") parser.add_argument("--config", default="configs/config.yaml", help="Config YAML path") parser.add_argument( "--model", choices=["simple_cnn", "mobilenet_v2"], help="Override model (default: from config)" ) parser.add_argument("--epochs", type=int, help="Override number of epochs") args = parser.parse_args() with open(args.config) as f: cfg = yaml.safe_load(f) train(cfg, model_name_override=args.model, epochs_override=args.epochs)