Gokulnathan G
Deploy DR Grading v10 — FastAPI + Material Design 3 UI with runtime model download
648beff | """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} | |