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:
rpotter6298
2026-04-28 08:24:25 +02:00
parent 4dea45df78
commit 512ebd13b2
42 changed files with 4468 additions and 612 deletions
+6
View File
@@ -109,6 +109,12 @@ BACKBONES: Dict[str, BackboneSpec] = {
strip=_strip_efficientnet,
blocks=_blocks_efficientnet,
),
"resnet18": BackboneSpec(
ctor=models.resnet18,
weights_default=models.ResNet18_Weights.DEFAULT,
strip=_strip_resnet,
blocks=_blocks_resnet,
),
"resnet50": BackboneSpec(
ctor=models.resnet50,
weights_default=models.ResNet50_Weights.DEFAULT,
+43
View File
@@ -44,6 +44,37 @@ class ImageTransformConfig:
]
return transforms.Compose(ops)
def build_precache(self) -> transforms.Compose:
"""Deterministic prefix: PIL → resized CHW float32 in [0, 1].
Output is suitable for caching; per-batch ``build_postcache`` finishes
the pipeline (augment + normalize) on tensors.
"""
return transforms.Compose([
transforms.Resize(self.resize_size),
transforms.CenterCrop(self.crop_size),
transforms.ToTensor(),
])
def build_postcache(self) -> transforms.Compose:
"""Per-batch tail run on cached float32 [0, 1] CHW tensors.
Augmentations operate on tensors (torchvision v1 supports this for
Flip/Rotation/ColorJitter on tensor input). Normalize is applied last.
"""
ops = []
if self.augment:
if self.hflip:
ops.append(transforms.RandomHorizontalFlip())
if self.vflip:
ops.append(transforms.RandomVerticalFlip())
if self.rotation_deg:
ops.append(transforms.RandomRotation(self.rotation_deg))
if self.color_jitter:
ops.append(transforms.ColorJitter(*self.color_jitter))
ops.append(transforms.Normalize(mean=self.mean, std=self.std))
return transforms.Compose(ops)
def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig:
"""Build an ImageTransformConfig using the backbone's default normalisation stats."""
@@ -64,3 +95,15 @@ def build_backbone_transform(backbone_name: str, augment: bool = True) -> transf
def build_eval_transform(backbone_name: str) -> transforms.Compose:
"""Deterministic eval transform — no augmentation, backbone-matched normalisation."""
return build_backbone_transform(backbone_name, augment=False)
def build_split_transforms(
backbone_name: str, augment: bool = True
) -> tuple[transforms.Compose, transforms.Compose]:
"""Return (precache, postcache) transform pair for tensor-cached image towers.
precache : PIL → CHW float32 in [0, 1] (deterministic, run once at fill)
postcache : tensor → augmented + normalized tensor (run per batch)
"""
cfg = backbone_transform_config(backbone_name, augment=augment)
return cfg.build_precache(), cfg.build_postcache()
+212
View File
@@ -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