Add distributed server implementation and protocol definitions
- 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.
This commit is contained in:
@@ -0,0 +1,212 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user