Download classification_segmentation/train.py from Ishaank18/aifh: direct link, hf CLI and curl.
- Browser
- Download file 15.7 kB
-
https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/train.py
- Command line
-
hf download hf://Ishaank18/aifh/classification_segmentation/train.py
-
curl -L -o train.py https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/train.py
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() | |