Segmentering
Kunna skilja semantisk och instanssegmentering och träna en enkel U-Net.
Förkunskaper
Intuition
Segmentering klassificerar varje pixel. Tre varianter som ofta blandas ihop:
| Variant | Svarar på | Två bilar bredvid varandra |
|---|---|---|
| Semantisk | vilken klass är varje pixel? | ett sammanhängande «bil»-område |
| Instans | vilket objekt tillhör pixeln? | två separata bilar |
| Panoptisk | båda | två bilar + bakgrundsklasser |
U-Net är arkitekturen som dominerat sedan 2015, särskilt i medicinsk bildanalys. Formen är ett U:
nedsampling uppsampling
64 ─────────── hoppkoppling ───────────→ 64
↓ 128 ─────── hoppkoppling ──────→ 128 ↑
↓ 256 ──── hoppkoppling ─→ 256 ↑
↓ 512 ──────────────→ 512 ↑
Hoppkopplingarna är hela poängen. Nedsamplingen ger sammanhang (vad är det här?) men tappar rumslig upplösning. Hoppkopplingarna tar tillbaka de finkorniga detaljerna från motsvarande nivå — utan dem blir maskerna suddiga och kanterna oprecisa.
Formellt
Förlustfunktionen avgör mer än arkitekturen vid segmentering, eftersom klasserna nästan alltid är extremt obalanserade — en tumör kan vara 0,1 % av pixlarna.
| Förlust | Formel | När |
|---|---|---|
| Korsentropi | per pixel | balanserade klasser |
| Viktad korsentropi | klassvikter | måttlig obalans |
| Dice | stark obalans | |
| Dice + CE | summa | standardvalet i praktiken |
| Focal | nedviktar lätta pixlar | extrem obalans |
| Tversky | Dice med justerbar FP/FN-vikt | när missar kostar mer än falsklarm |
Ren korsentropi på 0,1 % positiva pixlar ger en modell som förutsäger «bakgrund» överallt och får 99,9 % pixelträffsäkerhet. Dice-förlusten straffar det direkt, eftersom den mäter överlappet med den positiva klassen.
Mått:
| Mått | Formel | Notera |
|---|---|---|
| IoU / Jaccard | strängare | |
| Dice / F1 | mildare; Dice ≥ IoU alltid | |
| mIoU | medel över klasser | standard för semantisk segmentering |
| Hausdorff-avstånd | största kantavvikelse | när kanten är det viktiga |
Rapportera per klass, inte bara medelvärdet. En mIoU på 0,78 kan dölja att den viktigaste klassen ligger på 0,31.
Praktiska detaljer som spelar stor roll:
| Detalj | Varför |
|---|---|
| Augmentering med elastisk deformation | U-Net:s originalartikel lyfte detta som avgörande för medicinska bilder |
| Patch-baserad träning | hela bilder ryms sällan i minnet |
| Överlappande patchar vid inferens | annars syns sömmar i masken |
| Klassvikter från faktisk pixelfördelning | inte gissade |
SAM (Segment Anything) ändrade läget 2023: en promptbar modell tränad på en miljard masker som segmenterar godtyckliga objekt utan finjustering. För många uppgifter är rätt utgångspunkt i dag att pröva SAM först och bara träna eget om det inte räcker.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
def block(cin, cut):
return nn.Sequential(
nn.Conv2d(cin, cut, 3, padding=1, bias=False), nn.BatchNorm2d(cut), nn.ReLU(inplace=True),
nn.Conv2d(cut, cut, 3, padding=1, bias=False), nn.BatchNorm2d(cut), nn.ReLU(inplace=True))
class UNet(nn.Module):
def __init__(self, in_kanaler=3, n_klasser=1, bas=64):
super().__init__()
self.ner1, self.ner2 = block(in_kanaler, bas), block(bas, bas * 2)
self.ner3, self.ner4 = block(bas * 2, bas * 4), block(bas * 4, bas * 8)
self.botten = block(bas * 8, bas * 16)
self.upp4 = nn.ConvTranspose2d(bas * 16, bas * 8, 2, stride=2)
self.upp3 = nn.ConvTranspose2d(bas * 8, bas * 4, 2, stride=2)
self.upp2 = nn.ConvTranspose2d(bas * 4, bas * 2, 2, stride=2)
self.upp1 = nn.ConvTranspose2d(bas * 2, bas, 2, stride=2)
self.dek4, self.dek3 = block(bas * 16, bas * 8), block(bas * 8, bas * 4)
self.dek2, self.dek1 = block(bas * 4, bas * 2), block(bas * 2, bas)
self.ut = nn.Conv2d(bas, n_klasser, 1)
self.pool = nn.MaxPool2d(2)
def forward(self, x):
h1 = self.ner1(x)
h2 = self.ner2(self.pool(h1))
h3 = self.ner3(self.pool(h2))
h4 = self.ner4(self.pool(h3))
b = self.botten(self.pool(h4))
d = self.dek4(torch.cat([self.upp4(b), h4], 1)) # hoppkoppling
d = self.dek3(torch.cat([self.upp3(d), h3], 1))
d = self.dek2(torch.cat([self.upp2(d), h2], 1))
d = self.dek1(torch.cat([self.upp1(d), h1], 1))
return self.ut(d)
def dice_forlust(logits, mal, eps=1.0):
p = torch.sigmoid(logits)
snitt = (p * mal).sum(dim=(2, 3))
return (1 - (2 * snitt + eps) / (p.sum(dim=(2, 3)) + mal.sum(dim=(2, 3)) + eps)).mean()
def kombinerad(logits, mal, w_dice=0.5):
return w_dice * dice_forlust(logits, mal) + (1 - w_dice) * F.binary_cross_entropy_with_logits(logits, mal)
# Varför ren korsentropi misslyckas vid stark obalans
mal = torch.zeros(1, 1, 256, 256); mal[0, 0, 120:130, 120:130] = 1.0 # 0,15 % positiva
allt_bakgrund = torch.full((1, 1, 256, 256), -10.0) # förutsäg alltid 0
print("CE :", round(float(F.binary_cross_entropy_with_logits(allt_bakgrund, mal)), 5))
print("Dice:", round(float(dice_forlust(allt_bakgrund, mal)), 5))
# CE : 0.00153 ← nästan noll: modellen "lyckas" genom att aldrig hitta något
# Dice: 0.99999 ← straffar det direkt
# Mät per klass, inte bara medelvärdet
def iou_per_klass(pred, mal, n_klasser):
ut = {}
for k in range(n_klasser):
p, m = (pred == k), (mal == k)
union = float((p | m).sum())
ut[k] = round(float((p & m).sum()) / union, 4) if union else None
return ut
# Överlappande patchar vid inferens — annars syns sömmar
def segmentera_stor_bild(modell, bild, patch=512, overlapp=64):
H, W = bild.shape[-2:]
steg = patch - overlapp
ut = torch.zeros(1, 1, H, W)
vikt = torch.zeros(1, 1, H, W)
for y in range(0, H, steg):
for x in range(0, W, steg):
y2, x2 = min(y + patch, H), min(x + patch, W)
with torch.no_grad():
p = torch.sigmoid(modell(bild[..., y:y2, x:x2]))
ut[..., y:y2, x:x2] += p
vikt[..., y:y2, x:x2] += 1
return ut / vikt.clamp_min(1)
Utskriften i mitten visar varför förlustvalet är viktigare än arkitekturen: en modell som aldrig hittar något får korsentropi 0,0015 — praktiskt taget perfekt enligt det måttet.
Behärskning innebär
- Skiljer semantisk, instans- och panoptisk segmentering
- Förklarar U-Net:s hoppkopplingar
- Väljer förlustfunktion vid obalans
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — U-Net: Convolutional Networks for Biomedical Image Segmentation — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Segment Anything — arXiv (öppen åtkomst; licens per artikel)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause