Spaces:
Sleeping
Sleeping
| """ | |
| 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} | |
| 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) | |