EyeQC / src /fov.py
rdharini2001's picture
EyeQC
6fe482b
Raw
History Blame Contribute Delete
8.02 kB
"""
Robust retinal Field-of-View (FOV) / fundus-circle detection.
Every downstream QC metric must be computed ONLY inside the illuminated retinal
disc, never over the black/again-textured camera surround. Real fundus images
are frequently off-centre, letter-boxed, or cropped so that the true circle is
truncated by the frame - so the disc centre is NOT the image centre.
Strategy
--------
1. Segment the foreground (fundus is bright/coloured vs a dark surround) with a
floor-OR-Otsu threshold on the per-pixel channel max, robust to dim images.
2. Keep the largest hole-filled connected component.
3. Fit the fundus circle by algebraic (Kasa) least squares to the *true* arc of
the boundary - boundary points lying on the image frame are truncation edges
and are excluded, which recovers the real centre/radius even when the disc is
heavily cropped. Falls back to min-enclosing circle / equivalent radius when
the arc is too small to fit.
4. The ROI mask is the fitted disc clipped to the frame (optionally intersected
with the foreground), giving a clean circular ROI with a correctly located
centre.
"""
from __future__ import annotations
import numpy as np
import cv2
def _largest_component(mask):
num, labels, stats, _ = cv2.connectedComponentsWithStats(mask.astype(np.uint8), 8)
if num <= 1:
return mask.astype(bool)
largest = 1 + int(np.argmax(stats[1:, cv2.CC_STAT_AREA]))
return labels == largest
def _fill_holes(mask):
filled = mask.astype(np.uint8) * 255
ff = filled.copy()
h, w = ff.shape
m = np.zeros((h + 2, w + 2), np.uint8)
cv2.floodFill(ff, m, (0, 0), 255)
return (filled | cv2.bitwise_not(ff)).astype(bool)
def _foreground(rgb):
h, w = rgb.shape[:2]
chmax = rgb.max(axis=2) # bright in any channel = tissue
blur = cv2.GaussianBlur(chmax, (0, 0), sigmaX=max(h, w) / 200.0)
otsu_t, _ = cv2.threshold(blur, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
thr = max(8, min(otsu_t * 0.5, 45))
mask = blur > thr
k = max(3, int(round(min(h, w) * 0.012)) | 1)
ker = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))
mask = cv2.morphologyEx(mask.astype(np.uint8), cv2.MORPH_OPEN, ker)
mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, ker)
mask = _largest_component(mask)
return _fill_holes(mask)
def _kasa_circle_fit(xs, ys):
"""Algebraic least-squares circle fit. Returns (cx, cy, r)."""
x = xs.astype(np.float64); y = ys.astype(np.float64)
A = np.stack([x, y, np.ones_like(x)], axis=1)
b = x ** 2 + y ** 2
sol, *_ = np.linalg.lstsq(A, b, rcond=None)
cx = sol[0] / 2.0; cy = sol[1] / 2.0
r = np.sqrt(max(sol[2] + cx ** 2 + cy ** 2, 1e-6))
return float(cx), float(cy), float(r)
def _extent_estimate(mask):
"""Robust circle estimate from foreground extents (handles truncation):
the least-truncated axis gives the diameter; centres from span midpoints."""
h, w = mask.shape
ys, xs = np.where(mask)
if xs.size == 0:
return None
BIG = w + h
row_min = np.full(h, BIG); row_max = np.full(h, -1)
np.minimum.at(row_min, ys, xs); np.maximum.at(row_max, ys, xs)
col_min = np.full(w, BIG); col_max = np.full(w, -1)
np.minimum.at(col_min, xs, ys); np.maximum.at(col_max, xs, ys)
widths = np.where(row_max >= 0, row_max - row_min + 1, 0)
heights = np.where(col_max >= 0, col_max - col_min + 1, 0)
hmax = int(widths.max()); vmax = int(heights.max())
wr = int(widths.argmax()); hc = int(heights.argmax())
cx = (row_min[wr] + row_max[wr]) / 2.0
cy = (col_min[hc] + col_max[hc]) / 2.0
r = 0.5 * max(hmax, vmax)
return float(cx), float(cy), float(r), float(hmax), float(vmax)
def _fit_circle_from_mask(mask, h, w):
"""Fit the fundus circle robustly. Extent estimate gives a truncation-proof
radius; Kasa arc-fit refines the centre when a clean arc is available."""
cnts, _ = cv2.findContours(mask.astype(np.uint8), cv2.RETR_EXTERNAL,
cv2.CHAIN_APPROX_NONE)
if not cnts:
return None
cnt = max(cnts, key=cv2.contourArea).reshape(-1, 2)
xs, ys = cnt[:, 0], cnt[:, 1]
m = max(2, int(round(min(h, w) * 0.01)))
on_frame = (xs <= m) | (xs >= w - 1 - m) | (ys <= m) | (ys >= h - 1 - m)
frame_frac = float(on_frame.mean())
ext = _extent_estimate(mask)
ex_cx, ex_cy, ex_r, hmax, vmax = ext
arc = ~on_frame
if arc.sum() >= 40 and arc.sum() >= 0.12 * len(xs):
kcx, kcy, kr = _kasa_circle_fit(xs[arc], ys[arc])
# accept Kasa only if consistent with the extent-based radius
if 0.8 * ex_r <= kr <= 1.35 * ex_r and \
np.hypot(kcx - ex_cx, kcy - ex_cy) <= 0.5 * ex_r:
return float(kcx), float(kcy), float(kr), frame_frac
return ex_cx, ex_cy, ex_r, frame_frac
def detect_fov(rgb):
"""Detect the retinal fundus circle. Returns a dict with mask/cx/cy/radius
and quality descriptors (coverage, centering, completeness, circularity)."""
h, w = rgb.shape[:2]
fg = _foreground(rgb)
if fg.sum() < 50:
return dict(mask=np.ones((h, w), bool), cx=w / 2, cy=h / 2,
radius=min(h, w) / 2, coverage=1.0, centering=0.0,
completeness=0.0, circularity=0.0, truncated=1.0)
fit = _fit_circle_from_mask(fg, h, w)
cx, cy, radius, frame_frac = fit
# Full-frame mode: the fundus fills the frame (a tight crop with little/no
# visible surround). Treat the whole tissue region as the ROI and don't
# penalise "completeness" for an unobservable circle edge.
full_frame = frame_frac > 0.85 and fg.mean() > 0.95
if full_frame:
ys, xs = np.where(fg)
cx, cy = float(xs.mean()), float(ys.mean())
radius = float(min(h, w) / 2.0)
yy, xx = np.mgrid[0:h, 0:w]
mask = fg
area = float(mask.sum())
return dict(mask=mask.astype(bool), cx=cx, cy=cy, radius=radius,
coverage=area / (h * w), centering=1.0, completeness=1.0,
circularity=1.0, truncated=0.0, full_frame=True)
# ROI = fitted disc clipped to frame, intersected with foreground to stay
# inside real tissue where the disc overshoots.
yy, xx = np.mgrid[0:h, 0:w]
disc = (xx - cx) ** 2 + (yy - cy) ** 2 <= radius ** 2
mask = disc & (fg | disc) # disc within frame
# guard: if disc barely overlaps fg (bad fit), use fg directly
if (mask & fg).sum() < 0.5 * fg.sum():
mask = fg
ys, xs = np.where(fg); cx, cy = xs.mean(), ys.mean()
radius = float(np.sqrt(fg.sum() / np.pi))
disc = (xx - cx) ** 2 + (yy - cy) ** 2 <= radius ** 2
area = float(mask.sum())
coverage = area / (h * w)
off = np.hypot(cx - w / 2, cy - h / 2)
centering = float(np.clip(1 - off / (radius + 1e-6), 0, 1))
# completeness: fraction of the fitted circle that lies inside the frame
theta = np.linspace(0, 2 * np.pi, 360, endpoint=False)
px = (cx + radius * 0.98 * np.cos(theta))
py = (cy + radius * 0.98 * np.sin(theta))
inside_frame = (px >= 0) & (px < w) & (py >= 0) & (py < h)
completeness = float(inside_frame.mean())
# circularity of the detected foreground (isoperimetric)
cnts, _ = cv2.findContours(fg.astype(np.uint8), cv2.RETR_EXTERNAL,
cv2.CHAIN_APPROX_SIMPLE)
peri = cv2.arcLength(max(cnts, key=cv2.contourArea), True) if cnts else 0.0
circ = float(np.clip(4 * np.pi * fg.sum() / (peri ** 2 + 1e-6), 0, 1))
return dict(mask=mask.astype(bool), cx=float(cx), cy=float(cy),
radius=float(radius), coverage=coverage, centering=centering,
completeness=completeness, circularity=circ,
truncated=float(frame_frac))
def roi_mask_uint8(rgb, fov=None):
"""0/255 ROI mask for external models (e.g. RRWNet preprocessing)."""
if fov is None:
fov = detect_fov(rgb)
return (fov["mask"].astype(np.uint8) * 255)