Download classification_segmentation/common.py from Ishaank18/aifh: direct link, hf CLI and curl.
- Browser
- Download file 25.6 kB
-
https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/common.py
- Command line
-
hf download hf://Ishaank18/aifh/classification_segmentation/common.py
-
curl -L -o common.py https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/common.py
25.6 kB
| """ | |
| common.py -- shared configuration, datasets, transforms and metric helpers | |
| for Part A (COVID-19 Radiography Database). | |
| Seed policy (assignment requirement): the random seed is the last two digits of | |
| the roll number. Roll number ends in 16 -> SEED = 16. The same seed is used | |
| for every split, sample selection, augmentation seed and reported audit subset. | |
| """ | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import os | |
| import random | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Callable, Dict, List, Optional, Sequence, Tuple | |
| import numpy as np | |
| # ---------------------------------------------------------------------------- | |
| # 0. Global configuration | |
| # ---------------------------------------------------------------------------- | |
| SEED: int = 16 # last two digits of the roll number | |
| SEED_TAG: str = f"seed{SEED:02d}" # -> "seed16", used in every output filename | |
| CLASSES: List[str] = ["COVID", "Lung_Opacity", "Normal", "Viral Pneumonia"] | |
| CLASS_TO_IDX: Dict[str, int] = {c: i for i, c in enumerate(CLASSES)} | |
| IDX_TO_CLASS: Dict[int, str] = {i: c for c, i in CLASS_TO_IDX.items()} | |
| NUM_CLASSES: int = len(CLASSES) | |
| # Folder names as they ship inside COVID-19_Radiography_Dataset/. | |
| # (Some mirrors use "Viral Pneumonia", some "Viral_Pneumonia" -- both accepted.) | |
| CLASS_DIR_ALIASES: Dict[str, List[str]] = { | |
| "COVID": ["COVID", "Covid", "COVID-19"], | |
| "Lung_Opacity": ["Lung_Opacity", "Lung Opacity"], | |
| "Normal": ["Normal"], | |
| "Viral Pneumonia": ["Viral Pneumonia", "Viral_Pneumonia"], | |
| } | |
| IMG_SIZE: int = 224 # 224 matches TorchXRayVision res224 weights and DenseNet121 | |
| SPLIT_FRACTIONS = (0.70, 0.15, 0.15) # train / val / test, stratified by class | |
| # Repository layout (this file lives in <ROOT>/classification_segmentation/) | |
| ROOT = Path(__file__).resolve().parents[1] | |
| DATASET_DIR = ROOT / "Dataset" / "covid19" | |
| SPLIT_DIR = ROOT / "splits" | |
| RESULTS_DIR = ROOT / "Results" | |
| METRICS_DIR = RESULTS_DIR / "metrics" | |
| PLOTS_DIR = RESULTS_DIR / "plots" | |
| OVERLAY_DIR = RESULTS_DIR / "overlays" | |
| WEIGHTS_DIR = RESULTS_DIR / "model_weights" | |
| CLS_SPLIT_CSV = SPLIT_DIR / f"covid_classification_{SEED_TAG}.csv" | |
| SEG_SPLIT_CSV = SPLIT_DIR / f"covid_segmentation_{SEED_TAG}.csv" | |
| for _d in (SPLIT_DIR, METRICS_DIR, PLOTS_DIR, OVERLAY_DIR, WEIGHTS_DIR): | |
| _d.mkdir(parents=True, exist_ok=True) | |
| # ---------------------------------------------------------------------------- | |
| # 1. Reproducibility | |
| # ---------------------------------------------------------------------------- | |
| def set_seed(seed: int = SEED, deterministic: bool = True) -> None: | |
| """Seed python, numpy and torch (CPU + CUDA). Call at the top of every script.""" | |
| os.environ["PYTHONHASHSEED"] = str(seed) | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| try: | |
| import torch | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed_all(seed) | |
| if deterministic: | |
| torch.backends.cudnn.deterministic = True | |
| torch.backends.cudnn.benchmark = False | |
| except ImportError: # torch not needed for the pure-EDA path | |
| pass | |
| def seeded_rng(offset: int = 0) -> np.random.Generator: | |
| """A dedicated numpy Generator so that 'seed-selected' samples are stable | |
| and independent of how many random calls the training loop made.""" | |
| return np.random.default_rng(SEED + offset) | |
| def worker_init_fn(worker_id: int) -> None: | |
| np.random.seed(SEED + worker_id) | |
| random.seed(SEED + worker_id) | |
| def get_device(): | |
| import torch | |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| # ---------------------------------------------------------------------------- | |
| # 2. Small IO helpers | |
| # ---------------------------------------------------------------------------- | |
| def save_json(obj, path: Path) -> None: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with open(path, "w") as f: | |
| json.dump(_jsonable(obj), f, indent=2) | |
| print(f"[saved] {path}") | |
| def _jsonable(o): | |
| if isinstance(o, dict): | |
| return {str(k): _jsonable(v) for k, v in o.items()} | |
| if isinstance(o, (list, tuple)): | |
| return [_jsonable(v) for v in o] | |
| if isinstance(o, (np.integer,)): | |
| return int(o) | |
| if isinstance(o, (np.floating,)): | |
| return float(o) | |
| if isinstance(o, np.ndarray): | |
| return o.tolist() | |
| if isinstance(o, Path): | |
| return str(o) | |
| return o | |
| def file_md5(path: Path, chunk: int = 1 << 20) -> str: | |
| h = hashlib.md5() | |
| with open(path, "rb") as f: | |
| while True: | |
| b = f.read(chunk) | |
| if not b: | |
| break | |
| h.update(b) | |
| return h.hexdigest() | |
| def progress(iterable, desc: str = "", total: Optional[int] = None): | |
| """tqdm if available, otherwise a light-weight fallback (no hard dependency).""" | |
| try: | |
| from tqdm.auto import tqdm | |
| return tqdm(iterable, desc=desc, total=total, leave=False) | |
| except ImportError: | |
| return iterable | |
| # ---------------------------------------------------------------------------- | |
| # 3. Image / mask loading + preprocessing | |
| # ---------------------------------------------------------------------------- | |
| # NOTE ON A DATASET QUIRK (documented in the report): | |
| # In COVID-19_Radiography_Dataset v3+ the images are 299x299 PNG while the | |
| # lung masks are 256x256 PNG. They cover the *same* field of view, so the two | |
| # are aligned after both are resized to a common IMG_SIZE. We therefore never | |
| # assume image.size == mask.size; we resize both explicitly. | |
| def load_image_gray(path, size: int = IMG_SIZE) -> np.ndarray: | |
| """Read as 8-bit grayscale, resize bilinearly, return float32 in [0, 1].""" | |
| from PIL import Image | |
| img = Image.open(path).convert("L") | |
| if size is not None and img.size != (size, size): | |
| img = img.resize((size, size), Image.BILINEAR) | |
| return np.asarray(img, dtype=np.float32) / 255.0 | |
| def load_mask_bin(path, size: int = IMG_SIZE, threshold: float = 0.5) -> np.ndarray: | |
| """Read a lung mask, resize with NEAREST (never interpolate labels), binarize.""" | |
| from PIL import Image | |
| m = Image.open(path).convert("L") | |
| if size is not None and m.size != (size, size): | |
| m = m.resize((size, size), Image.NEAREST) | |
| arr = np.asarray(m, dtype=np.float32) / 255.0 | |
| return (arr > threshold).astype(np.float32) | |
| def border_mask_image(img: np.ndarray, frac: float = 0.06) -> np.ndarray: | |
| """Ablation: blank an outer ring to suppress acquisition marks / burned-in | |
| annotations / L-R laterality letters that sit near the image border.""" | |
| out = img.copy() | |
| h, w = out.shape[:2] | |
| bh, bw = int(round(h * frac)), int(round(w * frac)) | |
| out[:bh, :] = 0.0 | |
| out[h - bh:, :] = 0.0 | |
| out[:, :bw] = 0.0 | |
| out[:, w - bw:] = 0.0 | |
| return out | |
| def lung_crop_image(img: np.ndarray, mask: Optional[np.ndarray], | |
| margin: float = 0.05) -> np.ndarray: | |
| """Ablation: crop to the lung-mask bounding box (+margin) then resize back. | |
| Falls back to the untouched image when the mask is missing or empty.""" | |
| import cv2 | |
| if mask is None or mask.sum() < 10: | |
| return img | |
| ys, xs = np.where(mask > 0.5) | |
| y0, y1 = ys.min(), ys.max() | |
| x0, x1 = xs.min(), xs.max() | |
| h, w = img.shape[:2] | |
| my, mx = int(round((y1 - y0) * margin)), int(round((x1 - x0) * margin)) | |
| y0, y1 = max(0, y0 - my), min(h - 1, y1 + my) | |
| x0, x1 = max(0, x0 - mx), min(w - 1, x1 + mx) | |
| crop = img[y0:y1 + 1, x0:x1 + 1] | |
| if crop.size == 0: | |
| return img | |
| return cv2.resize(crop, (w, h), interpolation=cv2.INTER_LINEAR) | |
| # ---------------------------------------------------------------------------- | |
| # 4. Augmentation policies | |
| # ---------------------------------------------------------------------------- | |
| # Deliberate choice, be ready to defend it in the viva: | |
| # * NO horizontal flip by default. A chest radiograph has a fixed laterality | |
| # (heart on the patient's left, gastric bubble left, burned-in "L"/"R" | |
| # markers). Mirroring creates anatomically impossible images and destroys | |
| # the marker cue we explicitly study in the border-mask ablation. | |
| # * Small affine + mild photometric jitter only, which mimics realistic | |
| # positioning and exposure variation between acquisitions. | |
| AUG_POLICIES = { | |
| "none": dict(rotate=0.0, translate=0.00, scale=(1.00, 1.00), jitter=0.00, erase=0.0), | |
| "weak": dict(rotate=5.0, translate=0.03, scale=(0.95, 1.00), jitter=0.08, erase=0.0), | |
| "default": dict(rotate=7.0, translate=0.05, scale=(0.90, 1.00), jitter=0.12, erase=0.0), | |
| "strong": dict(rotate=15.0, translate=0.10, scale=(0.80, 1.00), jitter=0.30, erase=0.25), | |
| } | |
| def build_train_augment(policy: str = "default", img_size: int = IMG_SIZE): | |
| """Return a torchvision transform operating on a 1-channel float tensor.""" | |
| import torch | |
| try: | |
| from torchvision.transforms import v2 # torchvision >= 0.15 | |
| except ImportError: # older torchvision | |
| from torchvision import transforms as v2 | |
| p = AUG_POLICIES[policy] | |
| ops = [] | |
| if p["rotate"] > 0 or p["translate"] > 0 or p["scale"] != (1.0, 1.0): | |
| ops.append( | |
| v2.RandomAffine( | |
| degrees=p["rotate"], | |
| translate=(p["translate"], p["translate"]), | |
| scale=p["scale"], | |
| interpolation=v2.InterpolationMode.BILINEAR, | |
| ) | |
| ) | |
| if p["jitter"] > 0: | |
| ops.append(v2.ColorJitter(brightness=p["jitter"], contrast=p["jitter"])) | |
| if p["erase"] > 0: | |
| ops.append(v2.RandomErasing(p=p["erase"], scale=(0.02, 0.10), value=0.0)) | |
| return v2.Compose(ops) if ops else torch.nn.Identity() | |
| # ---------------------------------------------------------------------------- | |
| # 5. Model-specific intensity normalisation | |
| # ---------------------------------------------------------------------------- | |
| IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) | |
| IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) | |
| # IMPORTANT ORDERING CONSTRAINT | |
| # ---------------------------- | |
| # torchvision's photometric transforms (ColorJitter, RandomErasing, ...) assume | |
| # float images live in [0, 1] and CLAMP their output to that range. Applying | |
| # them *after* model-specific normalisation therefore destroys the signal: | |
| # a TorchXRayVision tensor spanning [-1024, 1024] collapses to [0, 1], and an | |
| # ImageNet-normalised tensor spanning roughly [-2.1, 2.6] does the same. | |
| # So the pipeline is strictly: | |
| # load -> [0,1] -> ablation -> augment (still [0,1]) -> normalise_for_model | |
| def normalize_for_model(x01, model_kind: str, channels: str = "auto"): | |
| """x01: torch tensor (1, H, W) in [0, 1]. Returns the model's input tensor. | |
| model_kind : "xrv" -> 1 channel rescaled to [-1024, 1024]; identical to | |
| xrv.datasets.normalize(x, 255) for 8-bit input | |
| "student" -> 3 channels (grayscale replicated) + ImageNet stats, | |
| or 1 channel when channels == "gray" | |
| channels : "auto" | "gray" | "rgb" (drives the grayscale-vs-RGB ablation) | |
| """ | |
| import torch | |
| if x01.ndim == 2: | |
| x01 = x01.unsqueeze(0) | |
| x01 = x01.float() | |
| if model_kind == "xrv": | |
| return (x01 * 2.0 - 1.0) * 1024.0 | |
| if channels != "gray": | |
| t = x01.repeat(3, 1, 1) if x01.shape[0] == 1 else x01 | |
| mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1) | |
| std = torch.tensor(IMAGENET_STD).view(3, 1, 1) | |
| return (t - mean) / std | |
| return (x01 - 0.449) / 0.226 # mean/std of the ImageNet grayscale average | |
| def to_model_tensor(img01: np.ndarray, model_kind: str, channels: str = "auto"): | |
| """Convenience wrapper for the no-augmentation paths (evaluation, Grad-CAM): | |
| numpy HxW in [0,1] -> normalised model input.""" | |
| import torch | |
| x = torch.from_numpy(np.asarray(img01, dtype=np.float32)) | |
| return normalize_for_model(x.unsqueeze(0), model_kind, channels) | |
| def denormalize_for_display(t) -> np.ndarray: | |
| """Turn any model input tensor back into a HxW array in [0,1] for plotting.""" | |
| import torch | |
| a = t.detach().cpu().float() | |
| if a.ndim == 4: | |
| a = a[0] | |
| if a.shape[0] == 3: | |
| mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1) | |
| std = torch.tensor(IMAGENET_STD).view(3, 1, 1) | |
| a = a * std + mean | |
| a = a.mean(0) | |
| else: | |
| a = a[0] | |
| if a.abs().max() > 10: # xrv range | |
| a = (a / 1024.0 + 1.0) / 2.0 | |
| else: | |
| a = a * 0.226 + 0.449 | |
| return a.clamp(0, 1).numpy() | |
| # ---------------------------------------------------------------------------- | |
| # 6. Datasets | |
| # ---------------------------------------------------------------------------- | |
| class SampleSpec: | |
| """One row of a split CSV.""" | |
| image_id: str | |
| image_path: str | |
| mask_path: str | |
| label: str | |
| split: str | |
| def _read_split(csv_path: Path, split: str) -> "pd.DataFrame": | |
| import pandas as pd | |
| if not Path(csv_path).exists(): | |
| raise FileNotFoundError( | |
| f"Split file {csv_path} not found. Run pre_process.py first." | |
| ) | |
| df = pd.read_csv(csv_path) | |
| if split != "all": | |
| df = df[df["split"] == split].reset_index(drop=True) | |
| return df | |
| class CovidClassificationDataset: | |
| """Returns (tensor, label_idx, image_id). Torch Dataset (duck-typed to avoid | |
| importing torch at module import time).""" | |
| def __init__(self, csv_path: Path, split: str, model_kind: str, | |
| img_size: int = IMG_SIZE, augment: str = "none", | |
| ablation: str = "none", channels: str = "auto"): | |
| import torch.utils.data # noqa: F401 (ensures torch present) | |
| self.df = _read_split(csv_path, split) | |
| self.split = split | |
| self.model_kind = model_kind | |
| self.img_size = img_size | |
| self.ablation = ablation | |
| self.channels = channels | |
| self.augment = build_train_augment(augment, img_size) if augment != "none" else None | |
| # ablations that need the lung mask at load time | |
| self.needs_mask = ablation == "lung_crop" | |
| def __len__(self) -> int: | |
| return len(self.df) | |
| def __getitem__(self, i: int): | |
| import torch | |
| row = self.df.iloc[i] | |
| img = load_image_gray(row["image_path"], self.img_size) | |
| if self.ablation == "lung_crop": | |
| mp = row.get("mask_path", "") | |
| mask = load_mask_bin(mp, self.img_size) if isinstance(mp, str) and mp and Path(mp).exists() else None | |
| img = lung_crop_image(img, mask) | |
| elif self.ablation == "border_mask": | |
| img = border_mask_image(img, frac=0.06) | |
| # augment while the tensor is still in [0, 1] -- see the ordering note | |
| # above normalize_for_model(); doing it afterwards would clamp away the | |
| # model-specific normalisation. | |
| x01 = torch.from_numpy(img.astype(np.float32)).unsqueeze(0) # (1, H, W) | |
| if self.augment is not None: | |
| x01 = self.augment(x01).clamp(0.0, 1.0) | |
| ch = self.channels | |
| if self.ablation == "gray_input": | |
| ch = "gray" | |
| x = normalize_for_model(x01, self.model_kind, channels=ch) | |
| y = CLASS_TO_IDX[row["label"]] | |
| return x, y, str(row["image_id"]) | |
| class CovidSegmentationDataset: | |
| """Returns (image_tensor, mask_tensor[1,H,W], image_id, label).""" | |
| def __init__(self, csv_path: Path, split: str, img_size: int = IMG_SIZE, | |
| augment: str = "none", in_channels: int = 1): | |
| self.df = _read_split(csv_path, split) | |
| self.df = self.df[self.df["has_mask"] == True].reset_index(drop=True) # noqa: E712 | |
| self.img_size = img_size | |
| self.in_channels = in_channels | |
| self.aug_policy = augment | |
| def __len__(self) -> int: | |
| return len(self.df) | |
| def __getitem__(self, i: int): | |
| import torch | |
| row = self.df.iloc[i] | |
| img = load_image_gray(row["image_path"], self.img_size) | |
| mask = load_mask_bin(row["mask_path"], self.img_size) | |
| if self.aug_policy != "none": | |
| img, mask = _joint_augment(img, mask, self.aug_policy) | |
| x = torch.from_numpy(img[None, ...].astype(np.float32)) | |
| if self.in_channels == 3: | |
| x = x.repeat(3, 1, 1) | |
| m = torch.from_numpy(mask[None, ...].astype(np.float32)) | |
| return x, m, str(row["image_id"]), str(row["label"]) | |
| def _joint_augment(img: np.ndarray, mask: np.ndarray, policy: str): | |
| """Geometric augmentation applied identically to image and mask | |
| (bilinear for the image, nearest for the mask).""" | |
| import cv2 | |
| p = AUG_POLICIES[policy] | |
| h, w = img.shape | |
| ang = np.random.uniform(-p["rotate"], p["rotate"]) if p["rotate"] else 0.0 | |
| sc = np.random.uniform(*p["scale"]) if p["scale"] != (1.0, 1.0) else 1.0 | |
| tx = np.random.uniform(-p["translate"], p["translate"]) * w | |
| ty = np.random.uniform(-p["translate"], p["translate"]) * h | |
| M = cv2.getRotationMatrix2D((w / 2, h / 2), ang, sc) | |
| M[0, 2] += tx | |
| M[1, 2] += ty | |
| img = cv2.warpAffine(img, M, (w, h), flags=cv2.INTER_LINEAR, borderValue=0.0) | |
| mask = cv2.warpAffine(mask, M, (w, h), flags=cv2.INTER_NEAREST, borderValue=0.0) | |
| if p["jitter"]: | |
| g = 1.0 + np.random.uniform(-p["jitter"], p["jitter"]) | |
| b = np.random.uniform(-p["jitter"], p["jitter"]) * 0.5 | |
| img = np.clip(img * g + b, 0.0, 1.0) | |
| return img.astype(np.float32), (mask > 0.5).astype(np.float32) | |
| def make_loader(dataset, batch_size: int, shuffle: bool, num_workers: int = 2, | |
| sampler=None): | |
| import torch | |
| from torch.utils.data import DataLoader | |
| g = torch.Generator() | |
| g.manual_seed(SEED) | |
| return DataLoader( | |
| dataset, | |
| batch_size=batch_size, | |
| shuffle=(shuffle and sampler is None), | |
| sampler=sampler, | |
| num_workers=num_workers, | |
| pin_memory=torch.cuda.is_available(), | |
| worker_init_fn=worker_init_fn, | |
| generator=g, | |
| drop_last=False, | |
| ) | |
| def class_weights_from_split(csv_path: Path, split: str = "train"): | |
| """Inverse-frequency weights, w_c = N / (K * n_c), normalised to mean 1.""" | |
| import torch | |
| df = _read_split(csv_path, split) | |
| counts = np.array([max(1, int((df["label"] == c).sum())) for c in CLASSES], dtype=np.float64) | |
| w = counts.sum() / (len(CLASSES) * counts) | |
| w = w / w.mean() | |
| return torch.tensor(w, dtype=torch.float32), counts.astype(int) | |
| # ---------------------------------------------------------------------------- | |
| # 7. Metrics | |
| # ---------------------------------------------------------------------------- | |
| def classification_metrics(y_true: np.ndarray, y_pred: np.ndarray, | |
| y_prob: Optional[np.ndarray] = None) -> Dict: | |
| from sklearn.metrics import (accuracy_score, balanced_accuracy_score, | |
| classification_report, confusion_matrix, | |
| f1_score, roc_auc_score) | |
| out: Dict = {} | |
| out["accuracy"] = float(accuracy_score(y_true, y_pred)) | |
| out["balanced_accuracy"] = float(balanced_accuracy_score(y_true, y_pred)) | |
| out["macro_f1"] = float(f1_score(y_true, y_pred, average="macro", zero_division=0)) | |
| out["weighted_f1"] = float(f1_score(y_true, y_pred, average="weighted", zero_division=0)) | |
| out["confusion_matrix"] = confusion_matrix( | |
| y_true, y_pred, labels=list(range(NUM_CLASSES))).tolist() | |
| rep = classification_report( | |
| y_true, y_pred, labels=list(range(NUM_CLASSES)), | |
| target_names=CLASSES, output_dict=True, zero_division=0) | |
| out["per_class"] = { | |
| c: {"precision": rep[c]["precision"], "recall": rep[c]["recall"], | |
| "f1": rep[c]["f1-score"], "support": int(rep[c]["support"])} | |
| for c in CLASSES | |
| } | |
| if y_prob is not None: | |
| try: | |
| out["auroc_macro_ovr"] = float( | |
| roc_auc_score(y_true, y_prob, multi_class="ovr", average="macro")) | |
| out["auroc_weighted_ovr"] = float( | |
| roc_auc_score(y_true, y_prob, multi_class="ovr", average="weighted")) | |
| per = {} | |
| for i, c in enumerate(CLASSES): | |
| yb = (np.asarray(y_true) == i).astype(int) | |
| per[c] = float(roc_auc_score(yb, y_prob[:, i])) if yb.sum() else float("nan") | |
| out["auroc_per_class"] = per | |
| except Exception as e: # e.g. a class absent from the split | |
| out["auroc_error"] = str(e) | |
| return out | |
| def segmentation_metrics(pred: np.ndarray, gt: np.ndarray, eps: float = 1e-7) -> Dict: | |
| """pred/gt: binary HxW arrays.""" | |
| p = (pred > 0.5).astype(np.float64) | |
| g = (gt > 0.5).astype(np.float64) | |
| tp = float((p * g).sum()) | |
| fp = float((p * (1 - g)).sum()) | |
| fn = float(((1 - p) * g).sum()) | |
| tn = float(((1 - p) * (1 - g)).sum()) | |
| return { | |
| "dice": (2 * tp) / (2 * tp + fp + fn + eps), | |
| "iou": tp / (tp + fp + fn + eps), | |
| "pixel_accuracy": (tp + tn) / (tp + tn + fp + fn + eps), | |
| "sensitivity": tp / (tp + fn + eps), # recall on lung pixels | |
| "specificity": tn / (tn + fp + eps), | |
| "precision": tp / (tp + fp + eps), | |
| } | |
| def boundary_metrics(pred: np.ndarray, gt: np.ndarray, tol: int = 2) -> Dict: | |
| """Boundary-sensitive metrics: symmetric HD95, mean surface distance and the | |
| boundary F1 (BF) score at a `tol`-pixel tolerance. These catch the case where | |
| Dice looks fine but the pleural/diaphragm edge is systematically wrong.""" | |
| from scipy.ndimage import binary_erosion, distance_transform_edt | |
| p = (pred > 0.5) | |
| g = (gt > 0.5) | |
| if p.sum() == 0 or g.sum() == 0: | |
| return {"hd95": float("nan"), "assd": float("nan"), "boundary_f1": 0.0} | |
| pb = p ^ binary_erosion(p) | |
| gb = g ^ binary_erosion(g) | |
| dt_g = distance_transform_edt(~gb) | |
| dt_p = distance_transform_edt(~pb) | |
| d_pg = dt_g[pb] # pred boundary -> gt boundary | |
| d_gp = dt_p[gb] # gt boundary -> pred boundary | |
| if d_pg.size == 0 or d_gp.size == 0: | |
| return {"hd95": float("nan"), "assd": float("nan"), "boundary_f1": 0.0} | |
| hd95 = float(max(np.percentile(d_pg, 95), np.percentile(d_gp, 95))) | |
| assd = float((d_pg.mean() + d_gp.mean()) / 2.0) | |
| prec = float((d_pg <= tol).mean()) | |
| rec = float((d_gp <= tol).mean()) | |
| bf = 0.0 if (prec + rec) == 0 else 2 * prec * rec / (prec + rec) | |
| return {"hd95": hd95, "assd": assd, "boundary_f1": bf, | |
| "boundary_precision": prec, "boundary_recall": rec, "tolerance_px": tol} | |
| # ---------------------------------------------------------------------------- | |
| # 8. Plot helpers | |
| # ---------------------------------------------------------------------------- | |
| def _mpl(): | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| return plt | |
| def plot_confusion_matrix(cm, title: str, out_path: Path, normalize: bool = False): | |
| plt = _mpl() | |
| cm = np.asarray(cm, dtype=np.float64) | |
| if normalize: | |
| cm = cm / np.clip(cm.sum(axis=1, keepdims=True), 1e-9, None) | |
| fig, ax = plt.subplots(figsize=(6.2, 5.4)) | |
| im = ax.imshow(cm, cmap="Blues") | |
| ax.set_xticks(range(NUM_CLASSES), CLASSES, rotation=35, ha="right") | |
| ax.set_yticks(range(NUM_CLASSES), CLASSES) | |
| ax.set_xlabel("Predicted") | |
| ax.set_ylabel("True") | |
| ax.set_title(title) | |
| fmt = "{:.2f}" if normalize else "{:.0f}" | |
| thr = cm.max() / 2 if cm.max() else 0.5 | |
| for i in range(NUM_CLASSES): | |
| for j in range(NUM_CLASSES): | |
| ax.text(j, i, fmt.format(cm[i, j]), ha="center", va="center", | |
| color="white" if cm[i, j] > thr else "black", fontsize=9) | |
| fig.colorbar(im, ax=ax, fraction=0.046) | |
| fig.tight_layout() | |
| Path(out_path).parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(out_path, dpi=160) | |
| plt.close(fig) | |
| print(f"[saved] {out_path}") | |
| def overlay_mask_on_image(img01: np.ndarray, mask: np.ndarray, | |
| color=(1.0, 0.25, 0.25), alpha: float = 0.35) -> np.ndarray: | |
| rgb = np.stack([img01] * 3, axis=-1) | |
| m = (mask > 0.5) | |
| for c in range(3): | |
| rgb[..., c] = np.where(m, (1 - alpha) * rgb[..., c] + alpha * color[c], rgb[..., c]) | |
| return np.clip(rgb, 0, 1) | |
| def overlay_heatmap(img01: np.ndarray, cam: np.ndarray, alpha: float = 0.4) -> np.ndarray: | |
| plt = _mpl() | |
| cmap = plt.get_cmap("jet") | |
| c = cam - cam.min() | |
| c = c / (c.max() + 1e-8) | |
| heat = cmap(c)[..., :3] | |
| rgb = np.stack([img01] * 3, axis=-1) | |
| return np.clip((1 - alpha) * rgb + alpha * heat, 0, 1) | |
| # ---------------------------------------------------------------------------- | |
| # 9. Version-compatibility shims (torch AMP API, matplotlib boxplot API) | |
| # ---------------------------------------------------------------------------- | |
| def make_grad_scaler(enabled: bool): | |
| import torch | |
| try: | |
| return torch.amp.GradScaler("cuda", enabled=enabled) # torch >= 2.4 | |
| except (AttributeError, TypeError): | |
| return torch.cuda.amp.GradScaler(enabled=enabled) # older torch | |
| def autocast(enabled: bool): | |
| import torch | |
| try: | |
| return torch.amp.autocast("cuda", enabled=enabled) # torch >= 2.4 | |
| except (AttributeError, TypeError): | |
| return torch.cuda.amp.autocast(enabled=enabled) | |
| def boxplot(ax, data, labels, **kw): | |
| """matplotlib renamed `labels` to `tick_labels` in 3.9.""" | |
| try: | |
| return ax.boxplot(data, tick_labels=labels, **kw) | |
| except TypeError: | |
| return ax.boxplot(data, labels=labels, **kw) | |