AI4 / preprocessing.py
Gokulnathan G
Deploy DR Grading v10 — FastAPI + Material Design 3 UI with runtime model download
648beff
Raw
History Blame Contribute Delete
7.34 kB
"""Image preprocessing and Grad-CAM utilities."""
import cv2
import numpy as np
import torch
import torch.nn.functional as F
from typing import List, Tuple, Optional, Dict
from dataclasses import dataclass
_clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
def apply_clahe_lab(rgb):
lab = cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB)
lab[:, :, 0] = _clahe.apply(lab[:, :, 0])
return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)
def retinal_mask(rgb):
gray = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY)
_, m = cv2.threshold(gray, 15, 255, cv2.THRESH_BINARY)
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))
m = cv2.morphologyEx(m, cv2.MORPH_CLOSE, kernel)
return cv2.morphologyEx(m, cv2.MORPH_OPEN, kernel)
def crop_retinal_disc(rgb, pad=10):
m = retinal_mask(rgb)
contours, _ = cv2.findContours(m, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return rgb
x, y, w, h = cv2.boundingRect(max(contours, key=cv2.contourArea))
x, y = max(0, x - pad), max(0, y - pad)
x2, y2 = min(rgb.shape[1], x + w + 2 * pad), min(rgb.shape[0], y + h + 2 * pad)
return rgb[y:y2, x:x2]
def _retinal_binary_mask(rgb):
if rgb.ndim != 3:
return np.ones(rgb.shape[:2], dtype=np.float32)
g = rgb[..., 1].astype(np.uint8)
_, m = cv2.threshold(g, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
if m.mean() < 30:
return np.ones(rgb.shape[:2], dtype=np.float32)
k = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (9, 9))
m = cv2.morphologyEx(m, cv2.MORPH_CLOSE, k)
m = cv2.morphologyEx(m, cv2.MORPH_OPEN, k)
m = cv2.erode(m, k, iterations=1)
m = cv2.GaussianBlur(m.astype(np.float32) / 255.0, (15, 15), 0)
return np.clip(m, 0.0, 1.0).astype(np.float32)
def retinal_soft_mask(h, w, margin_frac=0.015):
cy, cx = h / 2.0, w / 2.0
r_outer = min(h, w) * (0.5 - margin_frac)
r_inner = r_outer * 0.94
Y, X = np.mgrid[0:h, 0:w].astype(np.float32)
dist = np.sqrt((X - cx) ** 2 + (Y - cy) ** 2)
return np.clip((r_outer - dist) / max(r_outer - r_inner, 1.0), 0.0, 1.0).astype(np.float32)
def get_cam_target_layer(m):
bb = getattr(m, "backbone", None) or m
if hasattr(bb, "blocks"):
blks = bb.blocks
for idx_b, idx_s in [(-2, -1), (-2, None), (-1, -1), (-1, None)]:
try:
layer = blks[idx_b] if idx_s is None else blks[idx_b][idx_s]
if isinstance(layer, torch.nn.Module):
return [layer]
except (TypeError, IndexError):
continue
if hasattr(bb, "conv_head"):
return [bb.conv_head]
children = list(bb.children())
return [children[-1]] if children else [list(m.children())[-1]]
def tta_gradcam(cam_obj, inp, target_cls):
from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget
tgt = [ClassifierOutputTarget(target_cls)]
try:
c0 = cam_obj(input_tensor=inp, targets=tgt)[0]
c_h = cam_obj(input_tensor=torch.flip(inp, dims=[-1]), targets=tgt)[0]
c_v = cam_obj(input_tensor=torch.flip(inp, dims=[-2]), targets=tgt)[0]
return (c0 * 0.50 + np.fliplr(c_h) * 0.35 + np.flipud(c_v) * 0.15)
except Exception:
return cam_obj(input_tensor=inp, targets=tgt)[0]
def postprocess_cam(raw_cam, rgb_u8):
h, w = rgb_u8.shape[:2]
cam = cv2.resize(raw_cam.astype(np.float32), (w, h))
cam = cam * _retinal_binary_mask(rgb_u8)
vals = cam[cam > 0]
if vals.size > 10:
p2, p98 = float(np.percentile(vals, 2)), float(np.percentile(vals, 98))
cam = np.clip(cam, p2, p98)
cam = (cam - p2) / (p98 - p2 + 1e-8)
else:
cam = cam / (cam.max() + 1e-8)
cam_u8 = (np.clip(cam, 0, 1) * 255).astype(np.uint8)
cam_u8 = cv2.bilateralFilter(cam_u8, d=7, sigmaColor=45, sigmaSpace=7)
return np.power(cam_u8.astype(np.float32) / 255.0, 0.80).astype(np.float32)
def overlay_cam(rgb, cam, alpha=0.65):
rgb = rgb.astype(np.float32)
cam = np.clip(cam, 0.0, 1.0)
cam = cam * _retinal_binary_mask(rgb.astype(np.uint8))
if cam.max() > 1e-6:
cam = cam / cam.max()
cam_u8 = (cam * 255).astype(np.uint8)
heat = cv2.applyColorMap(cam_u8, cv2.COLORMAP_JET)
heat = cv2.cvtColor(heat, cv2.COLOR_BGR2RGB).astype(np.float32)
a_pix = alpha * np.power(cam, 0.7)[..., None]
blended = rgb * (1.0 - a_pix) + heat * a_pix
return np.clip(blended, 0, 255).astype(np.uint8)
def segment_lesions(cam, min_area_frac=0.0008):
h, w = cam.shape[:2]
min_area = max(40, int(h * w * min_area_frac))
cam_u8 = (np.clip(cam, 0.0, 1.0) * 255).astype(np.uint8)
otsu_thr, _ = cv2.threshold(cam_u8, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
nz = cam_u8[cam_u8 > 0]
p70 = float(np.percentile(nz, 70)) if nz.size > 0 else 128.0
thr = max(float(otsu_thr), p70)
binary = (cam_u8 >= thr).astype(np.uint8) * 255
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))
binary = cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel)
n_lab, labels, stats_, _ = cv2.connectedComponentsWithStats(binary, connectivity=8)
regions = []
for i in range(1, n_lab):
x, y, ww, hh, area = stats_[i]
if area < min_area:
continue
peak = float(cam[labels == i].max())
if peak >= 0.85: tier = "high"
elif peak >= 0.65: tier = "moderate"
elif peak >= 0.45: tier = "mild"
else: continue
regions.append({"bbox": (int(x), int(y), int(ww), int(hh)), "peak": peak, "tier": tier})
regions.sort(key=lambda r: r["peak"], reverse=True)
return regions[:8]
def draw_lesion_boxes(rgb, regions):
TIER_CLR = {"high": (255, 215, 0), "moderate": (255, 82, 82), "mild": (80, 165, 255)}
canvas = rgb.copy()
for r in regions:
x, y, w, h = r["bbox"]
clr = TIER_CLR[r["tier"]]
cv2.rectangle(canvas, (x, y), (x + w, y + h), clr, 2, cv2.LINE_AA)
label = f"{r['tier'][0].upper()} {r['peak']:.2f}"
(lw, lh), _ = cv2.getTextSize(label, cv2.FONT_HERSHEY_SIMPLEX, 0.42, 1)
ly = max(lh + 4, y - 2)
cv2.rectangle(canvas, (x, ly - lh - 4), (x + lw + 6, ly), clr, -1)
cv2.putText(canvas, label, (x + 3, ly - 3), cv2.FONT_HERSHEY_SIMPLEX, 0.42, (0, 0, 0), 1, cv2.LINE_AA)
return canvas
# Fundus heuristics
STRONG_HEUR = {"red_dominance": 0.60, "disc_coverage": 0.35, "edge_density": 0.010}
WEAK_HEUR = {"red_dominance": 0.35, "disc_coverage": 0.15, "edge_density": 0.008}
def fundus_heuristics(rgb):
r, g, b = rgb[..., 0].astype(np.float32), rgb[..., 1].astype(np.float32), rgb[..., 2].astype(np.float32)
bright = (r > 15) | (g > 15) | (b > 15)
n_bright = int(bright.sum())
if n_bright > 100:
red_dominance = float(((r > g) & (r > b))[bright].mean())
else:
red_dominance = float(((r > g) & (r > b)).mean())
mask = retinal_mask(rgb)
disc_coverage = float(mask.mean() / 255.0)
gray = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY)
gx = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=3)
gy = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=3)
mag = np.sqrt(gx**2 + gy**2)
edge_density = float(mag[bright].mean() / 255.0) if n_bright > 100 else float(mag.mean() / 255.0)
return {"red_dominance": red_dominance, "disc_coverage": disc_coverage, "edge_density": edge_density}