""" 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 /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 # ---------------------------------------------------------------------------- @dataclass 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)