Ishaank18's picture
Upload via upload_to_hf.py
d0518d9 verified
Raw History Blame Contribute Delete
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
# ----------------------------------------------------------------------------
@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)