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