Files
hypertower/v3/classes/seg_cnn.py
T
2026-04-14 19:42:16 +02:00

522 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Segmentation-map CNN for glaucoma grading.
Trains a CNN on combined disc/cup segmentation maps — pixel values
0 = background, 1 = disc (rim only), 2 = cup
— instead of raw RGB fundus images, forcing the model to learn purely
from optic nerve head geometry (CDR, rim width, cup location, etc.).
Two segmentation sources are supported:
gt rasterise expert contour/mask annotations directly (pure NumPy/PIL,
no CUDA — safe in DataLoader worker processes)
unet run a trained UNetSegmenter on the raw fundus image
Usage (import from training script):
from v3.classes.seg_cnn import SegMapRecord, SegMapDataset, SegCNN, seg_map_to_tensor
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import List, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from PIL import Image, ImageDraw
from PIL.Image import Resampling
from torch.utils.data import Dataset
from torchvision import models, transforms
from tqdm import tqdm
# ---------------------------------------------------------------------------
# Data record
# ---------------------------------------------------------------------------
@dataclass
class SegMapRecord:
"""One labelled eye sample for the seg-map CNN."""
sample_id: str
image_path: Path # original fundus image (used by unet mode)
annotation_disc: Path # contour (.txt) or mask (.bmp/.png)
annotation_cup: Path
annotation_type_disc: str # "contour" or "mask"
annotation_type_cup: str
patient_id: int # for group-CV: keep both eyes of a patient together
eye: str # "OD" or "OS"
label: int # 0 = Normal, 1 = Glaucoma
# ---------------------------------------------------------------------------
# Seg-map utilities
# ---------------------------------------------------------------------------
def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
"""
Combine binary disc and cup masks into a 3-class label map.
Returns a uint8 array with values:
0 — background
1 — optic disc rim (disc but not cup)
2 — optic cup
"""
disc = (disc_mask > 0).astype(np.uint8)
cup = (cup_mask > 0).astype(np.uint8)
cup = (cup & disc) # structural prior: cup must be inside disc
seg = disc + cup # 0, 1 (rim), or 2 (cup)
return seg.astype(np.uint8)
def crop_to_disc(seg_map: np.ndarray) -> np.ndarray:
"""
Crop a seg map tightly to the disc bounding box.
The disc is anywhere seg_map > 0 (i.e. rim or cup).
Returns the original array unchanged if no disc is found.
"""
rows = np.any(seg_map > 0, axis=1)
cols = np.any(seg_map > 0, axis=0)
if not rows.any():
return seg_map
r0, r1 = int(np.argmax(rows)), int(len(rows) - 1 - np.argmax(rows[::-1]))
c0, c1 = int(np.argmax(cols)), int(len(cols) - 1 - np.argmax(cols[::-1]))
return seg_map[r0:r1 + 1, c0:c1 + 1]
def seg_map_to_tensor(
seg_map: np.ndarray,
channels: int,
target_size: int,
) -> torch.Tensor:
"""
Convert an (H, W) seg map with values {0, 1, 2} to a float tensor.
channels=1 → (1, H, W) float in [0, 1] (values 0/0.5/1.0)
channels=3 → (3, H, W) one-hot binary channels [bg, disc_rim, cup]
"""
pil = Image.fromarray(seg_map.astype(np.uint8), mode="L")
pil = pil.resize((target_size, target_size), Image.NEAREST)
seg = np.array(pil, dtype=np.uint8)
if channels == 1:
arr = seg.astype(np.float32) / 2.0 # {0, 0.5, 1.0}
return torch.from_numpy(arr).unsqueeze(0)
if channels == 3:
bg = (seg == 0).astype(np.float32)
disc_rim = (seg == 1).astype(np.float32)
cup = (seg == 2).astype(np.float32)
return torch.from_numpy(np.stack([bg, disc_rim, cup], axis=0))
raise ValueError(f"channels must be 1 or 3, got {channels}")
# ---------------------------------------------------------------------------
# GT mask loading (pure NumPy / PIL — no CUDA, safe in DataLoader workers)
# ---------------------------------------------------------------------------
def _load_contour(path: Path) -> np.ndarray:
"""Load x,y contour pairs from a whitespace- or comma-delimited text file."""
for delimiter in (",", None):
try:
arr = np.loadtxt(str(path), delimiter=delimiter, comments="#", dtype=np.float32)
if arr.size > 0:
break
except Exception:
arr = np.zeros((0, 2), dtype=np.float32)
if arr.size == 0 or arr.ndim == 1:
return np.zeros((0, 2), dtype=np.float32)
if arr.shape[1] < 2:
return np.zeros((0, 2), dtype=np.float32)
return arr[:, :2]
def _contour_to_mask(
coords: np.ndarray, image_size: Tuple[int, int], target_size: int
) -> np.ndarray:
"""
Rasterise a polygon defined by (x, y) coords into a binary mask.
image_size is the (width, height) of the original fundus image — the
coordinate space the contour was annotated in. The mask is drawn at
that resolution then resized to target_size, matching UNetSegmenter's
behaviour and avoiding off-canvas clipping.
"""
if coords is None or len(coords) < 3:
return np.zeros((target_size, target_size), dtype=np.uint8)
points = [tuple(map(float, pt)) for pt in coords]
img = Image.new("L", image_size, 0)
ImageDraw.Draw(img).polygon(points, outline=1, fill=1)
img = img.resize((target_size, target_size), Resampling.NEAREST)
return (np.array(img, dtype=np.uint8) > 0).astype(np.uint8)
def _extract_masks_from_image(
mask_path: Path, target_size: int
) -> Tuple[np.ndarray, np.ndarray]:
"""
Extract disc and cup binary masks from a segmentation image file.
Handles both grayscale label images (e.g. REFUGE .bmp) and
RGB colour-coded masks. Returns (disc_mask, cup_mask) both at
target_size × target_size.
"""
raw = Image.open(mask_path)
arr = np.array(raw)
if arr.ndim == 2:
# Grayscale: identify background from edge statistics
edges = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]])
bg_val = int(np.argmax(np.bincount(edges.astype(np.int64).clip(0, 255), minlength=256)))
disc_arr = (arr != bg_val).astype(np.uint8)
vals = np.unique(arr)
non_bg = vals[vals != bg_val]
cup_arr: np.ndarray
if non_bg.size > 1:
cup_val = int(non_bg.min())
cup_arr = (arr == cup_val).astype(np.uint8)
else:
cup_arr = np.zeros_like(disc_arr, dtype=np.uint8)
else:
img_rgb = raw.convert("RGB")
arr = np.array(img_rgb)
h, w, c = arr.shape
edges_rgb = np.concatenate(
[arr[0], arr[-1], arr[:, 0], arr[:, -1]], axis=0
)
edge_colors, edge_counts = np.unique(edges_rgb.reshape(-1, c), axis=0, return_counts=True)
bg_color = edge_colors[int(np.argmax(edge_counts))]
colors, counts = np.unique(arr.reshape(-1, c), axis=0, return_counts=True)
not_bg = np.any(colors != bg_color.reshape(1, -1), axis=1)
colors, counts = colors[not_bg], counts[not_bg]
disc_arr = np.zeros((h, w), dtype=np.uint8)
cup_arr = np.zeros((h, w), dtype=np.uint8)
if colors.shape[0] >= 1:
order = np.argsort(-counts)
disc_color = colors[order[0]]
disc_arr[np.all(arr == disc_color, axis=-1)] = 1
if colors.shape[0] >= 2:
cup_color = colors[order[1]]
cup_arr[np.all(arr == cup_color, axis=-1)] = 1
# Resize to target_size with nearest-neighbour to preserve binary values
def _resize(m: np.ndarray) -> np.ndarray:
pil = Image.fromarray((m > 0).astype(np.uint8) * 255)
pil = pil.resize((target_size, target_size), Resampling.NEAREST)
return (np.array(pil) > 0).astype(np.uint8)
return _resize(disc_arr), _resize(cup_arr)
def load_gt_masks(rec: "SegMapRecord", target_size: int) -> Tuple[np.ndarray, np.ndarray]:
"""
Load GT disc + cup masks for one record.
Handles annotation_type "contour" (x,y text file) and "mask" (image file).
Returns (disc_mask, cup_mask) as uint8 arrays of shape (target_size, target_size).
"""
disc_mask: Optional[np.ndarray] = None
cup_mask: Optional[np.ndarray] = None
# Get original image size so contour coordinates are drawn in the right space
with Image.open(rec.image_path) as _img:
image_size = _img.size # (width, height)
# ---- Disc ----
if rec.annotation_type_disc == "mask":
disc_mask, cup_from_disc = _extract_masks_from_image(rec.annotation_disc, target_size)
if cup_from_disc.any():
cup_mask = cup_from_disc
else: # contour
coords = _load_contour(rec.annotation_disc)
disc_mask = _contour_to_mask(coords, image_size, target_size)
# ---- Cup ----
if cup_mask is None:
if rec.annotation_type_cup == "mask":
_, cup_from_cup = _extract_masks_from_image(rec.annotation_cup, target_size)
cup_mask = cup_from_cup
else: # contour
coords = _load_contour(rec.annotation_cup)
cup_mask = _contour_to_mask(coords, image_size, target_size)
if disc_mask is None:
disc_mask = np.zeros((target_size, target_size), dtype=np.uint8)
if cup_mask is None:
cup_mask = np.zeros((target_size, target_size), dtype=np.uint8)
# Structural prior: cup must lie within disc
cup_mask = (cup_mask > 0) & (disc_mask > 0)
return disc_mask.astype(np.uint8), cup_mask.astype(np.uint8)
# ---------------------------------------------------------------------------
# U-Net fine-tuning dataset
# ---------------------------------------------------------------------------
class UNetFineTuneDataset(Dataset):
"""
Loads (image_tensor, mask_tensor) pairs for fine-tuning the U-Net on
PAPILA GT annotations. Uses the same preprocessing as UNetSegmenter
so the fine-tuned weights are compatible with inference.
"""
def __init__(
self,
records: List[SegMapRecord],
target_size: int = 512,
normalize: str = "per_image",
) -> None:
self.records = records
self.target_size = target_size
self.normalize = normalize
self.to_tensor = transforms.ToTensor()
def __len__(self) -> int:
return len(self.records)
def _normalize(self, tensor: torch.Tensor) -> torch.Tensor:
if self.normalize == "per_image":
mean = tensor.mean(dim=(1, 2), keepdim=True)
std = tensor.std(dim=(1, 2), keepdim=True).clamp(min=1e-6)
return (tensor - mean) / std
if self.normalize == "imagenet":
mean = torch.tensor([0.485, 0.456, 0.406]).view(-1, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).view(-1, 1, 1)
return (tensor - mean) / std
return tensor
def __getitem__(self, idx: int):
rec = self.records[idx]
image = Image.open(rec.image_path).convert("RGB")
image = image.resize((self.target_size, self.target_size), Resampling.BILINEAR)
img_tensor = self._normalize(self.to_tensor(image))
disc_mask, cup_mask = load_gt_masks(rec, self.target_size)
mask_tensor = torch.from_numpy(
np.stack([disc_mask, cup_mask], axis=0).astype(np.float32)
)
return img_tensor, mask_tensor
# ---------------------------------------------------------------------------
# U-Net precomputation (run once per full record list, not per fold)
# ---------------------------------------------------------------------------
def precompute_unet_seg_maps(
records: List["SegMapRecord"],
segmenter,
threshold: float = 0.5,
) -> List[np.ndarray]:
"""
Run the U-Net on every record and return a list of combined seg maps.
Call this once before the CV loop and pass the results to each fold's
SegMapDataset via precomputed_seg_maps, so the U-Net isn't re-run per fold.
"""
to_tensor = transforms.ToTensor()
seg_maps = []
for rec in tqdm(records, desc="U-Net inference", unit="img", leave=False):
image = Image.open(rec.image_path).convert("RGB")
resized = segmenter.preprocess_image(image)
tensor = segmenter._normalize_tensor(
to_tensor(resized).to(segmenter.device)
).unsqueeze(0)
with torch.no_grad():
logits = segmenter.model(tensor)
probs = torch.sigmoid(logits)[0].cpu().numpy()
disc = (probs[0] > threshold).astype(np.uint8)
cup = (probs[1] > threshold).astype(np.uint8)
cup = (cup & disc)
seg_maps.append(_combine_masks(disc, cup.astype(np.uint8)))
return seg_maps
# ---------------------------------------------------------------------------
# Dataset
# ---------------------------------------------------------------------------
class SegMapDataset(Dataset):
"""
PyTorch Dataset that yields (seg_tensor, label) pairs.
Parameters
----------
records : list of SegMapRecord
target_size : CNN input spatial size (images are resized to this)
channels : 1 = single-channel label map; 3 = one-hot three channels
augment : apply random flips + rotation (for training set)
unet_segmenter : if provided, use U-Net predictions instead of GT masks;
must be a loaded UNetSegmenter with model weights set
unet_threshold : threshold for U-Net logit → binary mask
seg_target_size: resolution at which GT masks are rasterised (or U-Net
output size). Default 512 matches UNetSegmenter default.
crop_to_disc : crop the seg map tightly to the disc bounding box before
resizing to target_size (default True — eliminates the
background zeros that make up most of the full image)
"""
def __init__(
self,
records: List[SegMapRecord],
target_size: int = 224,
channels: int = 3,
augment: bool = False,
unet_segmenter=None,
unet_threshold: float = 0.5,
seg_target_size: int = 512,
crop_to_disc: bool = True,
precomputed_seg_maps: Optional[List[np.ndarray]] = None,
) -> None:
self.records = records
self.target_size = target_size
self.channels = channels
self.augment = augment
self.seg_target_size = seg_target_size
self.crop_to_disc = crop_to_disc
if precomputed_seg_maps is not None:
self._seg_maps = precomputed_seg_maps
elif unet_segmenter is not None:
self._seg_maps = precompute_unet_seg_maps(
records, unet_segmenter, unet_threshold
)
else:
self._seg_maps = None
# ------------------------------------------------------------------
def __len__(self) -> int:
return len(self.records)
# ------------------------------------------------------------------
def _augment(self, seg_map: np.ndarray) -> np.ndarray:
"""Random flips + 90° rotations (label-safe since NEAREST resize)."""
if np.random.rand() < 0.5:
seg_map = np.fliplr(seg_map)
if np.random.rand() < 0.5:
seg_map = np.flipud(seg_map)
k = np.random.randint(0, 4)
if k:
seg_map = np.rot90(seg_map, k=k)
return np.ascontiguousarray(seg_map)
# ------------------------------------------------------------------
def __getitem__(self, idx: int):
rec = self.records[idx]
if self._seg_maps is not None:
seg_map = self._seg_maps[idx]
else:
disc_mask, cup_mask = load_gt_masks(rec, self.seg_target_size)
seg_map = _combine_masks(disc_mask, cup_mask)
if self.crop_to_disc:
seg_map = crop_to_disc(seg_map)
if self.augment:
seg_map = self._augment(seg_map)
tensor = seg_map_to_tensor(seg_map, self.channels, self.target_size)
return tensor, rec.label
# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
class SegCNN(nn.Module):
"""
Pretrained CNN backbone adapted for segmentation-map input.
Parameters
----------
num_classes : output classes (2 for binary glaucoma grading)
backbone : "resnet18" | "resnet50" | "efficientnet_b0"
pretrained : initialise with ImageNet weights (recommended even for
non-RGB input — transfer generalises across domains)
in_channels : 1 (single label map) or 3 (one-hot channels)
dropout : dropout rate before the final classifier head
"""
def __init__(
self,
num_classes: int = 2,
backbone: str = "resnet18",
pretrained: bool = True,
in_channels: int = 3,
dropout: float = 0.3,
) -> None:
super().__init__()
weights_arg = "DEFAULT" if pretrained else None
if backbone == "resnet18":
base = models.resnet18(weights=weights_arg)
feat_dim = base.fc.in_features
base.fc = nn.Identity()
elif backbone == "resnet50":
base = models.resnet50(weights=weights_arg)
feat_dim = base.fc.in_features
base.fc = nn.Identity()
elif backbone == "efficientnet_b0":
base = models.efficientnet_b0(weights=weights_arg)
feat_dim = base.classifier[1].in_features
base.classifier = nn.Identity()
else:
raise ValueError(f"Unknown backbone: {backbone!r}")
# Adapt first conv layer if in_channels ≠ 3
if in_channels != 3:
first_conv = self._find_first_conv(base)
new_conv = nn.Conv2d(
in_channels,
first_conv.out_channels,
kernel_size=first_conv.kernel_size,
stride=first_conv.stride,
padding=first_conv.padding,
bias=first_conv.bias is not None,
)
if pretrained:
# Average pretrained RGB weights across channel dim
with torch.no_grad():
new_conv.weight.copy_(
first_conv.weight.mean(dim=1, keepdim=True).expand_as(new_conv.weight)
)
self._replace_first_conv(base, new_conv)
self.backbone = base
self.head = nn.Sequential(
nn.Dropout(p=dropout),
nn.Linear(feat_dim, num_classes),
)
# ------------------------------------------------------------------
@staticmethod
def _find_first_conv(module: nn.Module) -> nn.Conv2d:
for m in module.modules():
if isinstance(m, nn.Conv2d):
return m
raise RuntimeError("No Conv2d found in backbone")
@staticmethod
def _replace_first_conv(module: nn.Module, new_conv: nn.Conv2d) -> None:
"""Replace the first Conv2d in-place (handles resnet and efficientnet)."""
for name, child in module.named_children():
if isinstance(child, nn.Conv2d):
setattr(module, name, new_conv)
return
try:
SegCNN._replace_first_conv(child, new_conv)
return
except RuntimeError:
pass
raise RuntimeError("Could not replace first Conv2d")
# ------------------------------------------------------------------
def forward(self, x: torch.Tensor) -> torch.Tensor:
feats = self.backbone(x)
if feats.dim() > 2:
feats = feats.flatten(1)
return self.head(feats)