512ebd13b2
- Introduced `protocol.py` for shared data models used in server/client communication, including request and response schemas for registration, job submission, and status updates. - Implemented `server.py` to manage a SQLite job queue and client registry, handling job polling, status updates, and job completion. - Created a cheat sheet for server usage, detailing commands for starting the server, submitting jobs, and monitoring clients. - Added several experiment configuration files for various training setups, including geometry vector injections and baseline ensembles.
213 lines
8.6 KiB
Python
213 lines
8.6 KiB
Python
"""unet — REFUGE-trained UNet wrapper for v4.
|
|
|
|
Lean accessory module: model definition + a thin segmenter wrapper that handles
|
|
weight loading, preprocessing, inference, and fine-tuning.
|
|
|
|
Used by:
|
|
- GeometrySegEncoder tower (produces disc/cup seg maps as CNN input)
|
|
- (future) ImageEncoder cropping (locates disc bbox for image cropping)
|
|
|
|
The segmenter is intentionally domain-agnostic: it takes PIL images in and
|
|
returns binary (disc, cup) numpy masks. Fine-tuning consumes any DataLoader
|
|
yielding (image_tensor, mask_tensor) pairs — mask preparation (parsing GT
|
|
contour files, etc.) lives in the consumer.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
import torch
|
|
from PIL import Image
|
|
from PIL.Image import Resampling
|
|
from torch import nn
|
|
from torchvision import transforms
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Model
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class UNet(nn.Module):
|
|
def __init__(self, in_channels: int = 3, base_channels: int = 32, out_channels: int = 2):
|
|
super().__init__()
|
|
self.enc1 = self._block(in_channels, base_channels)
|
|
self.enc2 = self._block(base_channels, base_channels * 2)
|
|
self.enc3 = self._block(base_channels * 2, base_channels * 4)
|
|
self.enc4 = self._block(base_channels * 4, base_channels * 8)
|
|
|
|
self.pool = nn.MaxPool2d(2)
|
|
self.bottleneck = self._block(base_channels * 8, base_channels * 16)
|
|
|
|
self.up4 = nn.ConvTranspose2d(base_channels * 16, base_channels * 8, 2, stride=2)
|
|
self.dec4 = self._block(base_channels * 16, base_channels * 8)
|
|
self.up3 = nn.ConvTranspose2d(base_channels * 8, base_channels * 4, 2, stride=2)
|
|
self.dec3 = self._block(base_channels * 8, base_channels * 4)
|
|
self.up2 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 2, stride=2)
|
|
self.dec2 = self._block(base_channels * 4, base_channels * 2)
|
|
self.up1 = nn.ConvTranspose2d(base_channels * 2, base_channels, 2, stride=2)
|
|
self.dec1 = self._block(base_channels * 2, base_channels)
|
|
|
|
self.out_conv = nn.Conv2d(base_channels, out_channels, kernel_size=1)
|
|
|
|
@staticmethod
|
|
def _block(in_ch: int, out_ch: int) -> nn.Module:
|
|
return nn.Sequential(
|
|
nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),
|
|
nn.BatchNorm2d(out_ch),
|
|
nn.ReLU(inplace=True),
|
|
nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),
|
|
nn.BatchNorm2d(out_ch),
|
|
nn.ReLU(inplace=True),
|
|
)
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
e1 = self.enc1(x)
|
|
e2 = self.enc2(self.pool(e1))
|
|
e3 = self.enc3(self.pool(e2))
|
|
e4 = self.enc4(self.pool(e3))
|
|
b = self.bottleneck(self.pool(e4))
|
|
|
|
d4 = self.dec4(torch.cat([self.up4(b), e4], dim=1))
|
|
d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1))
|
|
d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))
|
|
d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))
|
|
return self.out_conv(d1)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Segmenter wrapper
|
|
# ---------------------------------------------------------------------------
|
|
|
|
_IMAGENET_MEAN = (0.485, 0.456, 0.406)
|
|
_IMAGENET_STD = (0.229, 0.224, 0.225)
|
|
|
|
|
|
class UNetSegmenter:
|
|
"""Wraps a UNet with preprocessing, weight loading, inference, and fine-tuning.
|
|
|
|
Parameters
|
|
----------
|
|
target_size : square resolution UNet operates at (default 512)
|
|
normalize : "per_image" | "imagenet" | "none"
|
|
device : torch device string; defaults to cuda if available
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
target_size: int = 512,
|
|
normalize: str = "per_image",
|
|
device: str | torch.device | None = None,
|
|
in_channels: int = 3,
|
|
base_channels: int = 32,
|
|
out_channels: int = 2,
|
|
):
|
|
self.target_size = target_size
|
|
self.normalize = normalize
|
|
self.device = (
|
|
torch.device(device) if device is not None
|
|
else torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
)
|
|
self.model = UNet(in_channels, base_channels, out_channels).to(self.device)
|
|
self._to_tensor = transforms.ToTensor()
|
|
|
|
# ── lifecycle ────────────────────────────────────────────────────────────
|
|
|
|
def to(self, device: str | torch.device) -> "UNetSegmenter":
|
|
self.device = torch.device(device)
|
|
self.model.to(self.device)
|
|
return self
|
|
|
|
def load_weights(self, path: str | Path) -> "UNetSegmenter":
|
|
"""Load a UNet checkpoint (raw state_dict or {'model': state_dict})."""
|
|
state = torch.load(Path(path), map_location=self.device, weights_only=False)
|
|
sd = state["model"] if isinstance(state, dict) and "model" in state else state
|
|
self.model.load_state_dict(sd)
|
|
self.model.eval()
|
|
return self
|
|
|
|
# ── preprocessing ────────────────────────────────────────────────────────
|
|
|
|
def _normalize_tensor(self, t: torch.Tensor) -> torch.Tensor:
|
|
if self.normalize == "per_image":
|
|
mean = t.mean(dim=(-2, -1), keepdim=True)
|
|
std = t.std (dim=(-2, -1), keepdim=True).clamp(min=1e-6)
|
|
return (t - mean) / std
|
|
if self.normalize == "imagenet":
|
|
mean = torch.tensor(_IMAGENET_MEAN, device=t.device).view(-1, 1, 1)
|
|
std = torch.tensor(_IMAGENET_STD, device=t.device).view(-1, 1, 1)
|
|
return (t - mean) / std
|
|
return t
|
|
|
|
def preprocess(self, image: Image.Image) -> torch.Tensor:
|
|
"""PIL image → normalized (C, H, W) tensor on segmenter device."""
|
|
resized = image.convert("RGB").resize(
|
|
(self.target_size, self.target_size), Resampling.BILINEAR
|
|
)
|
|
return self._normalize_tensor(self._to_tensor(resized).to(self.device))
|
|
|
|
# ── inference ────────────────────────────────────────────────────────────
|
|
|
|
@torch.no_grad()
|
|
def predict(
|
|
self,
|
|
image: Image.Image,
|
|
*,
|
|
threshold: float = 0.5,
|
|
tta: bool = False,
|
|
) -> tuple[np.ndarray, np.ndarray]:
|
|
"""Single image → (disc_mask, cup_mask) binary uint8 arrays at target_size.
|
|
|
|
cup_mask is restricted to disc area (cup ⊆ disc).
|
|
"""
|
|
self.model.eval()
|
|
x = self.preprocess(image).unsqueeze(0)
|
|
logits = self.model(x)
|
|
if tta:
|
|
log_h = torch.flip(self.model(torch.flip(x, dims=[3])), dims=[3])
|
|
log_v = torch.flip(self.model(torch.flip(x, dims=[2])), dims=[2])
|
|
logits = (logits + log_h + log_v) / 3.0
|
|
probs = torch.sigmoid(logits)[0].cpu().numpy()
|
|
disc = (probs[0] > threshold).astype(np.uint8)
|
|
cup = ((probs[1] > threshold) & (disc > 0)).astype(np.uint8)
|
|
return disc, cup
|
|
|
|
# ── fine-tuning ──────────────────────────────────────────────────────────
|
|
|
|
def finetune(
|
|
self,
|
|
dataloader,
|
|
*,
|
|
epochs: int = 10,
|
|
lr: float = 1e-5,
|
|
log_prefix: str = "[UNetSegmenter]",
|
|
) -> "UNetSegmenter":
|
|
"""Fine-tune on (image_tensor, mask_tensor) pairs.
|
|
|
|
image_tensor : (B, C, H, W) — already preprocessed (normalized)
|
|
mask_tensor : (B, 2, H, W) float32 — channel 0 disc, channel 1 cup
|
|
"""
|
|
import time
|
|
opt = torch.optim.Adam(self.model.parameters(), lr=lr)
|
|
crit = nn.BCEWithLogitsLoss()
|
|
for ep in range(1, epochs + 1):
|
|
self.model.train()
|
|
running, n_batches, t0 = 0.0, 0, time.time()
|
|
for img, mask in dataloader:
|
|
img, mask = img.to(self.device), mask.to(self.device)
|
|
opt.zero_grad()
|
|
loss = crit(self.model(img), mask)
|
|
loss.backward()
|
|
opt.step()
|
|
running += float(loss.item())
|
|
n_batches += 1
|
|
avg = running / max(n_batches, 1)
|
|
print(
|
|
f" {log_prefix} ep{ep:03d}/{epochs:03d} loss={avg:.4f} "
|
|
f"({time.time() - t0:.1f}s)",
|
|
flush=True,
|
|
)
|
|
self.model.eval()
|
|
return self
|