Ishaank18's picture
Upload via upload_to_hf.py
d0518d9 verified
Raw History Blame Contribute Delete
15.7 kB
"""
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()