zaidshabbir's picture
Upload 9 files
81d1ded verified
Raw
History Blame Contribute Delete
9.13 kB
"""
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)