""" train.py -- Part A training driver. Handles Task 1 (classification: `xrv` and `student`) and Task 2 (segmentation: `unet`) from the *same* frozen seed-16 split files. Examples -------- # Task 1, baseline runs python train.py --task classify --model xrv --epochs 20 python train.py --task classify --model student --epochs 25 --init-from imagenet # Task 1, the seeded ablation (identical for both models) python train.py --task classify --model xrv --ablation lung_crop --epochs 20 python train.py --task classify --model student --ablation lung_crop --epochs 25 # Task 2 python train.py --task segment --model unet --epochs 30 """ from __future__ import annotations import argparse import json import time from pathlib import Path import numpy as np import torch import torch.nn as nn from common import (CLASSES, make_grad_scaler, autocast, CLS_SPLIT_CSV, IMG_SIZE, METRICS_DIR, PLOTS_DIR, SEED, SEED_TAG, SEG_SPLIT_CSV, WEIGHTS_DIR, CovidClassificationDataset, CovidSegmentationDataset, _mpl, class_weights_from_split, get_device, make_loader, progress, save_json, segmentation_metrics, set_seed) from model import DiceBCELoss, FocalLoss, build_model, model_kind_for # --------------------------------------------------------------------------- def run_tag(args) -> str: """Deterministic, self-describing name used for weights / metrics / plots.""" if args.task == "classify": base = f"cls_{args.model}" if args.model == "student": base += f"_{args.init_from}" else: base = f"seg_{args.model}" if args.ablation != "none": base += f"_abl-{args.ablation}" if args.tag: base += f"_{args.tag}" return f"{base}_{SEED_TAG}" # --------------------------------------------------------------------------- # Classification # --------------------------------------------------------------------------- def train_classification(args): device = get_device() kind = model_kind_for(args.model) channels = "gray" if args.ablation == "gray_input" else "auto" aug_train = {"aug_strong": "strong", "aug_weak": "weak"}.get(args.ablation, args.augment) ds_tr = CovidClassificationDataset(CLS_SPLIT_CSV, "train", kind, args.img_size, augment=aug_train, ablation=args.ablation, channels=channels) ds_va = CovidClassificationDataset(CLS_SPLIT_CSV, "val", kind, args.img_size, augment="none", ablation=args.ablation, channels=channels) print(f"[data] train={len(ds_tr)} val={len(ds_va)} aug={aug_train} ablation={args.ablation}") weights, counts = class_weights_from_split(CLS_SPLIT_CSV, "train") print(f"[data] train class counts {dict(zip(CLASSES, counts.tolist()))}") print(f"[data] class weights {dict(zip(CLASSES, np.round(weights.numpy(), 3).tolist()))}") sampler = None if args.imbalance == "sampler": from torch.utils.data import WeightedRandomSampler per_sample = np.array([weights[CLASSES.index(l)] for l in ds_tr.df["label"]]) g = torch.Generator(); g.manual_seed(SEED) sampler = WeightedRandomSampler(torch.tensor(per_sample, dtype=torch.double), num_samples=len(ds_tr), replacement=True, generator=g) dl_tr = make_loader(ds_tr, args.bs, shuffle=True, num_workers=args.workers, sampler=sampler) dl_va = make_loader(ds_va, args.bs, shuffle=False, num_workers=args.workers) in_ch = 1 if (kind == "xrv" or channels == "gray") else 3 model = build_model(args.model, init_from=args.init_from, in_channels=in_ch, xrv_weights=args.xrv_weights).to(device) print(f"[model] {args.model}: {model.trainable_parameter_report()}") if args.imbalance == "class_weight": crit = nn.CrossEntropyLoss(weight=weights.to(device), label_smoothing=args.label_smoothing) elif args.imbalance == "focal": crit = FocalLoss(gamma=2.0, weight=weights.to(device)) else: crit = nn.CrossEntropyLoss(label_smoothing=args.label_smoothing) params = [p for p in model.parameters() if p.requires_grad] opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=args.wd) sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs, eta_min=args.lr * 0.02) scaler = make_grad_scaler(enabled=(device.type == "cuda" and args.amp)) from sklearn.metrics import balanced_accuracy_score, f1_score hist, best, bad_epochs = [], -1.0, 0 tag = run_tag(args) ckpt_path = WEIGHTS_DIR / f"{tag}.pt" t0 = time.time() for ep in range(1, args.epochs + 1): model.train() tl, n = 0.0, 0 for x, y, _ in progress(dl_tr, desc=f"ep{ep} train"): x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True) opt.zero_grad(set_to_none=True) with autocast(enabled=scaler.is_enabled()): loss = crit(model(x), y) scaler.scale(loss).backward() scaler.step(opt) scaler.update() tl += loss.item() * x.size(0); n += x.size(0) tr_loss = tl / max(1, n) model.eval() vl, m, yt, yp = 0.0, 0, [], [] with torch.no_grad(): for x, y, _ in progress(dl_va, desc=f"ep{ep} val"): x, y = x.to(device), y.to(device) logits = model(x) vl += crit(logits, y).item() * x.size(0); m += x.size(0) yt.append(y.cpu().numpy()); yp.append(logits.argmax(1).cpu().numpy()) yt, yp = np.concatenate(yt), np.concatenate(yp) va_loss = vl / max(1, m) macro_f1 = f1_score(yt, yp, average="macro", zero_division=0) bal_acc = balanced_accuracy_score(yt, yp) sched.step() hist.append({"epoch": ep, "train_loss": tr_loss, "val_loss": va_loss, "val_macro_f1": float(macro_f1), "val_balanced_acc": float(bal_acc), "lr": opt.param_groups[0]["lr"]}) print(f"[{tag}] ep {ep:03d} | train {tr_loss:.4f} | val {va_loss:.4f} " f"| macroF1 {macro_f1:.4f} | balAcc {bal_acc:.4f}") if macro_f1 > best + args.min_delta: best, bad_epochs = macro_f1, 0 torch.save({"model_state": model.state_dict(), "args": vars(args), "epoch": ep, "val_macro_f1": float(macro_f1), "classes": CLASSES, "seed": SEED, "tag": tag}, ckpt_path) print(f" -> new best, saved {ckpt_path.name}") else: bad_epochs += 1 if bad_epochs >= args.patience: print(f"[{tag}] early stop: no val macro-F1 gain > {args.min_delta} " f"for {args.patience} epochs") break summary = {"tag": tag, "task": "classify", "seed": SEED, "best_val_macro_f1": float(best), "epochs_run": len(hist), "minutes": round((time.time() - t0) / 60, 2), "config": vars(args), "history": hist, "train_class_counts": dict(zip(CLASSES, counts.tolist()))} save_json(summary, METRICS_DIR / f"train_{tag}.json") plot_curves(hist, ["train_loss", "val_loss"], ["val_macro_f1", "val_balanced_acc"], tag, PLOTS_DIR / f"curves_{tag}.png") print(f"[done] best val macro-F1 = {best:.4f} | weights: {ckpt_path}") # --------------------------------------------------------------------------- # Segmentation # --------------------------------------------------------------------------- def train_segmentation(args): device = get_device() in_ch = args.in_channels aug_train = {"aug_strong": "strong", "aug_weak": "weak"}.get(args.ablation, args.augment) ds_tr = CovidSegmentationDataset(SEG_SPLIT_CSV, "train", args.img_size, augment=aug_train, in_channels=in_ch) ds_va = CovidSegmentationDataset(SEG_SPLIT_CSV, "val", args.img_size, augment="none", in_channels=in_ch) print(f"[data] train={len(ds_tr)} val={len(ds_va)} aug={aug_train}") dl_tr = make_loader(ds_tr, args.bs, shuffle=True, num_workers=args.workers) dl_va = make_loader(ds_va, args.bs, shuffle=False, num_workers=args.workers) model = build_model("unet", in_channels=in_ch, base=args.base, bilinear=not args.transpose_up).to(device) n_par = sum(p.numel() for p in model.parameters()) print(f"[model] U-Net base={args.base} bilinear={not args.transpose_up} params={n_par/1e6:.2f}M") crit = DiceBCELoss(bce_weight=args.bce_weight) opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.wd) sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs, eta_min=args.lr * 0.02) scaler = make_grad_scaler(enabled=(device.type == "cuda" and args.amp)) hist, best, bad_epochs = [], -1.0, 0 tag = run_tag(args) ckpt_path = WEIGHTS_DIR / f"{tag}.pt" t0 = time.time() for ep in range(1, args.epochs + 1): model.train() tl, n = 0.0, 0 for x, m, _, _ in progress(dl_tr, desc=f"ep{ep} train"): x, m = x.to(device, non_blocking=True), m.to(device, non_blocking=True) opt.zero_grad(set_to_none=True) with autocast(enabled=scaler.is_enabled()): loss = crit(model(x), m) scaler.scale(loss).backward() scaler.step(opt); scaler.update() tl += loss.item() * x.size(0); n += x.size(0) tr_loss = tl / max(1, n) model.eval() vl, k, dices, ious = 0.0, 0, [], [] with torch.no_grad(): for x, m, _, _ in progress(dl_va, desc=f"ep{ep} val"): x, m = x.to(device), m.to(device) logits = model(x) vl += crit(logits, m).item() * x.size(0); k += x.size(0) pr = (torch.sigmoid(logits) > 0.5).float().cpu().numpy() gt = m.cpu().numpy() for i in range(pr.shape[0]): s = segmentation_metrics(pr[i, 0], gt[i, 0]) dices.append(s["dice"]); ious.append(s["iou"]) va_loss, dice, iou = vl / max(1, k), float(np.mean(dices)), float(np.mean(ious)) sched.step() hist.append({"epoch": ep, "train_loss": tr_loss, "val_loss": va_loss, "val_dice": dice, "val_iou": iou, "lr": opt.param_groups[0]["lr"]}) print(f"[{tag}] ep {ep:03d} | train {tr_loss:.4f} | val {va_loss:.4f} " f"| Dice {dice:.4f} | IoU {iou:.4f}") if dice > best + args.min_delta: best, bad_epochs = dice, 0 torch.save({"model_state": model.state_dict(), "args": vars(args), "epoch": ep, "val_dice": dice, "seed": SEED, "tag": tag}, ckpt_path) print(f" -> new best, saved {ckpt_path.name}") else: bad_epochs += 1 if bad_epochs >= args.patience: print(f"[{tag}] early stop after {ep} epochs") break summary = {"tag": tag, "task": "segment", "seed": SEED, "best_val_dice": float(best), "epochs_run": len(hist), "minutes": round((time.time() - t0) / 60, 2), "config": vars(args), "history": hist, "mask_preprocessing": { "resize_interpolation": "NEAREST (labels must not be interpolated)", "image_resize_interpolation": "BILINEAR", "binarization_threshold": 0.5, "target_size": args.img_size, "note": "source images are 299x299 and source masks 256x256; both are " "resized to the common target size, they share the same FOV", "cleaning": "none applied by default; empty masks are flagged in the audit " "and excluded via has_mask"}} save_json(summary, METRICS_DIR / f"train_{tag}.json") plot_curves(hist, ["train_loss", "val_loss"], ["val_dice", "val_iou"], tag, PLOTS_DIR / f"curves_{tag}.png") print(f"[done] best val Dice = {best:.4f} | weights: {ckpt_path}") # --------------------------------------------------------------------------- def plot_curves(hist, loss_keys, metric_keys, tag, out_path: Path): plt = _mpl() ep = [h["epoch"] for h in hist] fig, axes = plt.subplots(1, 2, figsize=(11, 4.2)) for k in loss_keys: axes[0].plot(ep, [h[k] for h in hist], marker="o", ms=3, label=k) axes[0].set_xlabel("epoch"); axes[0].set_ylabel("loss"); axes[0].legend() axes[0].set_title(f"{tag}: loss") for k in metric_keys: axes[1].plot(ep, [h[k] for h in hist], marker="o", ms=3, label=k) axes[1].set_xlabel("epoch"); axes[1].legend(); axes[1].set_title(f"{tag}: validation metrics") axes[1].grid(alpha=0.3) fig.tight_layout(); fig.savefig(out_path, dpi=160); plt.close(fig) print(f"[saved] {out_path}") def build_argparser(): ap = argparse.ArgumentParser(description="Part A training (seed 16)") ap.add_argument("--task", choices=["classify", "segment"], required=True) ap.add_argument("--model", default="xrv", choices=["xrv", "student", "unet"]) ap.add_argument("--ablation", default="none", choices=["none", "lung_crop", "border_mask", "gray_input", "aug_strong", "aug_weak"], help="the single clinically meaningful preprocessing change " "applied identically to both classification models") ap.add_argument("--epochs", type=int, default=20) ap.add_argument("--bs", type=int, default=32) ap.add_argument("--lr", type=float, default=None, help="default: 1e-3 (xrv head) / 3e-4 else") ap.add_argument("--wd", type=float, default=1e-4) ap.add_argument("--img-size", type=int, default=IMG_SIZE) ap.add_argument("--augment", default="default", choices=["none", "weak", "default", "strong"]) ap.add_argument("--imbalance", default="class_weight", choices=["none", "class_weight", "sampler", "focal"]) ap.add_argument("--label-smoothing", type=float, default=0.0) ap.add_argument("--init-from", default="imagenet", choices=["imagenet", "scratch"], help="Model 2 only: ImageNet (non-X-ray) init or random init") ap.add_argument("--xrv-weights", default="densenet121-res224-all") ap.add_argument("--base", type=int, default=32, help="U-Net base width") ap.add_argument("--in-channels", type=int, default=1, help="U-Net input channels") ap.add_argument("--transpose-up", action="store_true", help="U-Net: transposed conv upsampling") ap.add_argument("--bce-weight", type=float, default=0.5) ap.add_argument("--patience", type=int, default=6) ap.add_argument("--min-delta", type=float, default=1e-4) ap.add_argument("--workers", type=int, default=2) ap.add_argument("--amp", action="store_true", default=True) ap.add_argument("--tag", default="") return ap def main(): args = build_argparser().parse_args() if args.task == "segment": args.model = "unet" if args.lr is None: args.lr = 1e-3 if args.model == "xrv" else 3e-4 set_seed(SEED) print(json.dumps({"run": run_tag(args), "seed": SEED, **vars(args)}, indent=2, default=str)) if args.task == "classify": train_classification(args) else: train_segmentation(args) if __name__ == "__main__": main()