""" evaluate.py -- Part A evaluation on the FROZEN test split. Classification (Task 1) confusion matrix (raw + row-normalised), per-class precision/recall/F1, macro-F1, balanced accuracy, AUROC (per class, macro & weighted OvR), ROC curves, and an error analysis of >= 8 test errors per model with image identifier, true label, predicted label, confidence and quantitative descriptors that support the written explanation. Segmentation (Task 2) Dice, IoU, pixel accuracy, sensitivity, specificity, plus boundary-sensitive metrics (HD95, ASSD, boundary-F1 @2px), broken down by radiograph class, and >= 8 seed-selected overlays covering both successes and failures. Examples -------- python evaluate.py --task classify --ckpt ../Results/model_weights/cls_xrv_seed16.pt python evaluate.py --task segment --ckpt ../Results/model_weights/seg_unet_seed16.pt python evaluate.py --compare cls_xrv_seed16 cls_xrv_abl-lung_crop_seed16 \ cls_student_imagenet_seed16 cls_student_imagenet_abl-lung_crop_seed16 """ from __future__ import annotations import argparse import time from pathlib import Path import numpy as np import pandas as pd import torch from common import (CLASSES, boxplot, CLS_SPLIT_CSV, IDX_TO_CLASS, IMG_SIZE, METRICS_DIR, NUM_CLASSES, OVERLAY_DIR, PLOTS_DIR, SEED, SEED_TAG, SEG_SPLIT_CSV, WEIGHTS_DIR, CovidClassificationDataset, CovidSegmentationDataset, _mpl, boundary_metrics, classification_metrics, get_device, load_image_gray, load_mask_bin, make_loader, overlay_mask_on_image, progress, save_json, seeded_rng, segmentation_metrics, set_seed) from model import build_model, model_kind_for # --------------------------------------------------------------------------- def _torch_load(path): """torch>=2.6 defaults weights_only=True, which cannot restore the config dict we store alongside the weights; older torch has no such argument.""" try: return torch.load(path, map_location="cpu", weights_only=False) except TypeError: return torch.load(path, map_location="cpu") def load_checkpoint(ckpt_path: Path): ck = _torch_load(ckpt_path) cfg = ck["args"] task = cfg.get("task", "classify") name = cfg.get("model", "xrv") if task == "classify": kind = model_kind_for(name) channels = "gray" if cfg.get("ablation") == "gray_input" else "auto" in_ch = 1 if (kind == "xrv" or channels == "gray") else 3 model = build_model(name, init_from=cfg.get("init_from", "imagenet"), in_channels=in_ch, xrv_weights=cfg.get("xrv_weights")) else: model = build_model("unet", in_channels=cfg.get("in_channels", 1), base=cfg.get("base", 32), bilinear=not cfg.get("transpose_up", False)) model.load_state_dict(ck["model_state"]) model.eval() return model, ck, cfg # --------------------------------------------------------------------------- # Classification # --------------------------------------------------------------------------- @torch.no_grad() def predict_classification(model, cfg, split: str = "test", bs: int = 32, workers: int = 2): device = get_device() model.to(device).eval() kind = model_kind_for(cfg.get("model", "xrv")) channels = "gray" if cfg.get("ablation") == "gray_input" else "auto" ds = CovidClassificationDataset(CLS_SPLIT_CSV, split, kind, cfg.get("img_size", IMG_SIZE), augment="none", ablation=cfg.get("ablation", "none"), channels=channels) dl = make_loader(ds, bs, shuffle=False, num_workers=workers) ys, ps, ids, t0 = [], [], [], time.time() for x, y, idx in progress(dl, desc=f"predict[{split}]"): p = torch.softmax(model(x.to(device)), dim=1).cpu().numpy() ps.append(p); ys.append(y.numpy()); ids.extend(list(idx)) dt = time.time() - t0 return (np.concatenate(ys), np.concatenate(ps), ids, {"n_images": len(ids), "total_seconds": dt, "ms_per_image": 1000 * dt / max(1, len(ids))}) def plot_roc(y_true, y_prob, tag, out_path: Path): from sklearn.metrics import auc, roc_curve plt = _mpl() fig, ax = plt.subplots(figsize=(6.2, 5.2)) for i, c in enumerate(CLASSES): yb = (y_true == i).astype(int) if yb.sum() == 0: continue fpr, tpr, _ = roc_curve(yb, y_prob[:, i]) ax.plot(fpr, tpr, lw=1.8, label=f"{c} (AUC={auc(fpr, tpr):.3f})") ax.plot([0, 1], [0, 1], "k--", lw=0.8) ax.set_xlabel("False positive rate"); ax.set_ylabel("True positive rate") ax.set_title(f"ROC (one-vs-rest) -- {tag}") ax.legend(fontsize=8, loc="lower right"); ax.grid(alpha=0.3) fig.tight_layout(); fig.savefig(out_path, dpi=160); plt.close(fig) print(f"[saved] {out_path}") def describe_error(row, img, mask) -> str: """Quantitative, reproducible descriptors that back the written explanation. (The student still writes the clinical sentence -- these are the numbers it must be consistent with.)""" bits = [] border = np.concatenate([img[:14].ravel(), img[-14:].ravel(), img[:, :14].ravel(), img[:, -14:].ravel()]) bits.append(f"mean intensity {img.mean():.3f}") bits.append(f"contrast (std) {img.std():.3f}") bits.append(f"border brightness {border.mean():.3f}") if border.mean() > 0.35: bits.append("bright border/collimation or burned-in marker likely present") if img.mean() < 0.35: bits.append("under-exposed relative to the cohort mean") elif img.mean() > 0.65: bits.append("over-exposed relative to the cohort mean") if mask is not None: area = float(mask.mean()) bits.append(f"lung-mask area fraction {area:.3f}") if area < 0.18: bits.append("small lung field -> heavy consolidation or tight cropping") ys, xs = np.where(mask > 0.5) if len(ys): bits.append(f"lung bbox {xs.min()},{ys.min()}-{xs.max()},{ys.max()}") else: bits.append("no lung mask available for this image") return "; ".join(bits) def error_analysis(y_true, y_prob, ids, tag, n: int = 8, out_dir: Path = OVERLAY_DIR): """Pick the n most confident mistakes (most informative failures) and, if the seed subset asks for more diversity, top them up with seed-random errors.""" plt = _mpl() df = pd.read_csv(CLS_SPLIT_CSV) df = df[df["split"] == "test"] if df["image_id"].duplicated().any(): n = int(df["image_id"].duplicated().sum()) print(f"[warn] {n} duplicated image_id value(s) in the test split; " f"keeping the first occurrence of each for the error figure") df = df.drop_duplicates(subset="image_id", keep="first") df = df.set_index("image_id") y_pred = y_prob.argmax(1) conf = y_prob.max(1) err = np.where(y_pred != y_true)[0] if len(err) == 0: print(f"[{tag}] no test errors -- nothing to analyse") return pd.DataFrame() order = err[np.argsort(-conf[err])] chosen = list(order[:max(n // 2, 4)]) rng = seeded_rng(offset=7) rest = [i for i in order if i not in chosen] if rest: extra = rng.choice(len(rest), size=min(n - len(chosen), len(rest)), replace=False) chosen += [rest[i] for i in extra] chosen = chosen[:max(n, 8)] recs = [] ncol = 4 nrow = int(np.ceil(len(chosen) / ncol)) fig, axes = plt.subplots(nrow, ncol, figsize=(3.4 * ncol, 3.7 * nrow)) axes = np.atleast_1d(axes).ravel() for k, i in enumerate(chosen): iid = ids[i] row = df.loc[iid] if iid in df.index else None img = load_image_gray(row["image_path"], IMG_SIZE) if row is not None else np.zeros((IMG_SIZE, IMG_SIZE)) mask = None if row is not None and isinstance(row.get("mask_path"), str) and row["mask_path"] \ and Path(row["mask_path"]).exists(): mask = load_mask_bin(row["mask_path"], IMG_SIZE) rec = { "rank": k + 1, "image_id": iid, "true_label": IDX_TO_CLASS[int(y_true[i])], "predicted_label": IDX_TO_CLASS[int(y_pred[i])], "confidence": float(conf[i]), "prob_true_class": float(y_prob[i, int(y_true[i])]), "evidence": describe_error(row, img, mask), } recs.append(rec) axes[k].imshow(overlay_mask_on_image(img, mask) if mask is not None else img, cmap="gray") axes[k].set_title(f"{iid}\nT:{rec['true_label']} P:{rec['predicted_label']}\n" f"conf={rec['confidence']:.2f}", fontsize=8) axes[k].axis("off") for k in range(len(chosen), len(axes)): axes[k].axis("off") fig.suptitle(f"Test errors -- {tag} (seed={SEED})", fontsize=12) fig.tight_layout() out = out_dir / f"errors_{tag}.png" fig.savefig(out, dpi=150); plt.close(fig) print(f"[saved] {out}") edf = pd.DataFrame(recs) edf.to_csv(METRICS_DIR / f"errors_{tag}.csv", index=False) print(f"[saved] {METRICS_DIR / f'errors_{tag}.csv'}") print(edf[["image_id", "true_label", "predicted_label", "confidence"]].to_string(index=False)) return edf def evaluate_classification(args): model, ck, cfg = load_checkpoint(Path(args.ckpt)) tag = ck.get("tag", Path(args.ckpt).stem) y_true, y_prob, ids, timing = predict_classification(model, cfg, args.split, args.bs, args.workers) y_pred = y_prob.argmax(1) m = classification_metrics(y_true, y_pred, y_prob) m.update({"tag": tag, "split": args.split, "seed": SEED, "timing": timing, "checkpoint": str(args.ckpt), "config": cfg, "trainable_parameter_report": model.trainable_parameter_report()}) save_json(m, METRICS_DIR / f"eval_{tag}_{args.split}.json") from common import plot_confusion_matrix plot_confusion_matrix(m["confusion_matrix"], f"{tag} -- {args.split}", PLOTS_DIR / f"cm_{tag}_{args.split}.png") plot_confusion_matrix(m["confusion_matrix"], f"{tag} -- {args.split} (row-normalised)", PLOTS_DIR / f"cm_norm_{tag}_{args.split}.png", normalize=True) plot_roc(y_true, y_prob, tag, PLOTS_DIR / f"roc_{tag}_{args.split}.png") pd.DataFrame({"image_id": ids, "true": [IDX_TO_CLASS[i] for i in y_true], "pred": [IDX_TO_CLASS[i] for i in y_pred], "confidence": y_prob.max(1), **{f"p_{c}": y_prob[:, i] for i, c in enumerate(CLASSES)}} ).to_csv(METRICS_DIR / f"preds_{tag}_{args.split}.csv", index=False) if args.split == "test": error_analysis(y_true, y_prob, ids, tag, n=args.n_errors) print(f"\n=== {tag} / {args.split} ===") print(f" accuracy {m['accuracy']:.4f}") print(f" balanced accuracy {m['balanced_accuracy']:.4f}") print(f" macro-F1 {m['macro_f1']:.4f}") print(f" AUROC (macro OvR) {m.get('auroc_macro_ovr', float('nan')):.4f}") print(f" inference {timing['ms_per_image']:.2f} ms/image") for c in CLASSES: p = m["per_class"][c] print(f" {c:16s} P={p['precision']:.3f} R={p['recall']:.3f} " f"F1={p['f1']:.3f} n={p['support']}") # --------------------------------------------------------------------------- # Segmentation # --------------------------------------------------------------------------- @torch.no_grad() def evaluate_segmentation(args): device = get_device() model, ck, cfg = load_checkpoint(Path(args.ckpt)) tag = ck.get("tag", Path(args.ckpt).stem) model.to(device).eval() ds = CovidSegmentationDataset(SEG_SPLIT_CSV, args.split, cfg.get("img_size", IMG_SIZE), augment="none", in_channels=cfg.get("in_channels", 1)) dl = make_loader(ds, args.bs, shuffle=False, num_workers=args.workers) rows, t0 = [], time.time() for x, m, idxs, labels in progress(dl, desc=f"segment[{args.split}]"): pr = (torch.sigmoid(model(x.to(device))) > args.threshold).float().cpu().numpy() gt = m.numpy() for i in range(pr.shape[0]): s = segmentation_metrics(pr[i, 0], gt[i, 0]) b = boundary_metrics(pr[i, 0], gt[i, 0], tol=args.boundary_tol) rows.append({"image_id": idxs[i], "label": labels[i], **s, **b}) dt = time.time() - t0 df = pd.DataFrame(rows) num = df.select_dtypes(include=[np.number]) overall = {k: float(np.nanmean(num[k])) for k in num.columns} by_class = {c: {k: float(np.nanmean(g[k])) for k in num.columns} for c, g in df.groupby("label")} # is the difference across the four radiograph classes real? anova = {} try: from scipy.stats import f_oneway, kruskal groups = [g["dice"].dropna().values for _, g in df.groupby("label") if len(g) > 2] groups = [g for g in groups if len(g) > 2 and float(np.std(g)) > 1e-9] if len(groups) < 2: anova = {"note": "not enough non-degenerate class groups to test " "(a group had zero variance, e.g. an untrained model " "predicting a constant mask)"} elif len(groups) >= 2: f, p = f_oneway(*groups) h, pk = kruskal(*groups) anova = {"anova_F": float(f), "anova_p": float(p), "kruskal_H": float(h), "kruskal_p": float(pk)} except Exception as e: anova = {"error": str(e)} out = {"tag": tag, "split": args.split, "seed": SEED, "threshold": args.threshold, "n_images": int(len(df)), "overall": overall, "by_radiograph_class": by_class, "class_difference_test": anova, "checkpoint": str(args.ckpt), "config": cfg, "timing": {"total_seconds": dt, "ms_per_image": 1000 * dt / max(1, len(df))}} save_json(out, METRICS_DIR / f"eval_{tag}_{args.split}.json") df.to_csv(METRICS_DIR / f"per_image_{tag}_{args.split}.csv", index=False) print(f"\n=== {tag} / {args.split} ({len(df)} images) ===") for k in ("dice", "iou", "pixel_accuracy", "sensitivity", "specificity", "hd95", "assd", "boundary_f1"): if k in overall: print(f" {k:16s} {overall[k]:.4f}") print("\n Dice by radiograph class:") for c in CLASSES: if c in by_class: print(f" {c:16s} Dice={by_class[c]['dice']:.4f} IoU={by_class[c]['iou']:.4f} " f"HD95={by_class[c].get('hd95', float('nan')):.2f}") if anova and "anova_p" in anova: print(f" one-way ANOVA on Dice across classes: F={anova['anova_F']:.2f}, " f"p={anova['anova_p']:.2e} (Kruskal p={anova['kruskal_p']:.2e})") seg_overlays(model, cfg, df, tag, args) seg_boxplot(df, tag) @torch.no_grad() def seg_overlays(model, cfg, df: pd.DataFrame, tag: str, args, n: int = 8): """Seed-selected overlays: half seed-random, plus the worst cases so that both successes and failures are shown, as the assignment requires.""" plt = _mpl() device = get_device() rng = seeded_rng(offset=3) good = df.sort_values("dice", ascending=False) bad = df.sort_values("dice", ascending=True) n_rand = n // 2 rand_idx = rng.choice(len(df), size=min(n_rand, len(df)), replace=False) sel_ids = list(df.iloc[rand_idx]["image_id"]) + list(bad.head(n - n_rand)["image_id"]) sel_ids = list(dict.fromkeys(sel_ids))[:n] if len(sel_ids) < n: sel_ids += list(good.head(n - len(sel_ids))["image_id"]) split_df = pd.read_csv(SEG_SPLIT_CSV) split_df = split_df[split_df["split"] == args.split].set_index("image_id") per_img = df.set_index("image_id") ncol = 4 nrow = int(np.ceil(len(sel_ids) / ncol)) fig, axes = plt.subplots(nrow, ncol, figsize=(3.6 * ncol, 3.9 * nrow)) axes = np.atleast_1d(axes).ravel() in_ch = cfg.get("in_channels", 1) for k, iid in enumerate(sel_ids): r = split_df.loc[str(iid)] img = load_image_gray(r["image_path"], cfg.get("img_size", IMG_SIZE)) gt = load_mask_bin(r["mask_path"], cfg.get("img_size", IMG_SIZE)) x = torch.from_numpy(img[None, None].astype(np.float32)) if in_ch == 3: x = x.repeat(1, 3, 1, 1) pred = (torch.sigmoid(model(x.to(device)))[0, 0].cpu().numpy() > args.threshold) rgb = np.stack([img] * 3, -1) rgb[..., 1] = np.where(gt > 0.5, 0.5 * rgb[..., 1] + 0.5, rgb[..., 1]) # GT green rgb[..., 0] = np.where(pred, 0.5 * rgb[..., 0] + 0.5, rgb[..., 0]) # pred red d = float(per_img.loc[str(iid), "dice"]) axes[k].imshow(np.clip(rgb, 0, 1)) axes[k].set_title(f"{iid}\n{r['label']} | Dice={d:.3f}\n" f"HD95={per_img.loc[str(iid), 'hd95']:.1f}px", fontsize=8) axes[k].axis("off") for k in range(len(sel_ids), len(axes)): axes[k].axis("off") fig.suptitle(f"Lung segmentation overlays -- {tag} " f"(green = ground truth, red = prediction, seed={SEED})", fontsize=11) fig.tight_layout() out = OVERLAY_DIR / f"seg_overlays_{tag}_{args.split}.png" fig.savefig(out, dpi=150); plt.close(fig) print(f"[saved] {out}") def seg_boxplot(df: pd.DataFrame, tag: str): plt = _mpl() fig, axes = plt.subplots(1, 2, figsize=(11, 4.2)) boxplot(axes[0], [df[df.label == c]["dice"].dropna().values for c in CLASSES if c in set(df.label)], [c for c in CLASSES if c in set(df.label)], showfliers=False) axes[0].set_title("Dice by radiograph class"); axes[0].set_ylabel("Dice") plt.setp(axes[0].get_xticklabels(), rotation=20, ha="right") axes[1].hist(df["dice"].dropna(), bins=50, color="#55A868") axes[1].set_title("Dice distribution"); axes[1].set_xlabel("Dice") fig.tight_layout() out = PLOTS_DIR / f"seg_dice_by_class_{tag}.png" fig.savefig(out, dpi=160); plt.close(fig) print(f"[saved] {out}") # --------------------------------------------------------------------------- # Ablation / model comparison table # --------------------------------------------------------------------------- def compare_runs(tags, split: str = "test"): rows = [] for t in tags: p = METRICS_DIR / f"eval_{t}_{split}.json" if not p.exists(): print(f"[warn] missing {p}; run evaluate.py for '{t}' first") continue import json m = json.loads(p.read_text()) if "macro_f1" in m: rows.append({"run": t, "accuracy": m["accuracy"], "balanced_acc": m["balanced_accuracy"], "macro_F1": m["macro_f1"], "AUROC_macro": m.get("auroc_macro_ovr", float("nan")), **{f"F1_{c}": m["per_class"][c]["f1"] for c in CLASSES}, "ms_per_image": m["timing"]["ms_per_image"]}) else: o = m["overall"] rows.append({"run": t, "dice": o["dice"], "iou": o["iou"], "pixel_acc": o["pixel_accuracy"], "sensitivity": o["sensitivity"], "specificity": o["specificity"], "hd95": o.get("hd95"), "boundary_f1": o.get("boundary_f1")}) if not rows: return df = pd.DataFrame(rows).round(4) out = METRICS_DIR / f"comparison_{split}_{SEED_TAG}.csv" df.to_csv(out, index=False) print(f"\n=== Comparison ({split}) ===") print(df.to_string(index=False)) print(f"[saved] {out}") def main(): ap = argparse.ArgumentParser(description="Part A evaluation (seed 16)") ap.add_argument("--task", choices=["classify", "segment"]) ap.add_argument("--ckpt") ap.add_argument("--split", default="test", choices=["train", "val", "test"]) ap.add_argument("--bs", type=int, default=32) ap.add_argument("--workers", type=int, default=2) ap.add_argument("--threshold", type=float, default=0.5, help="segmentation binarisation") ap.add_argument("--boundary-tol", type=int, default=2) ap.add_argument("--n-errors", type=int, default=8) ap.add_argument("--compare", nargs="*", default=None, help="run tags to tabulate side by side") args = ap.parse_args() set_seed(SEED) if args.compare: compare_runs(args.compare, args.split) return if not args.task or not args.ckpt: ap.error("--task and --ckpt are required (or use --compare)") if args.task == "classify": evaluate_classification(args) else: evaluate_segmentation(args) if __name__ == "__main__": main()