Фаза 04 · урок 07
Семантическая сегментация — U-Net
Цель урока: Классификация выдаёт одну метку на изображение. Детекция выдаёт несколько рамок на изображение. Сегментация выдаёт одну метку на пиксель. Для входа размера H x W выход — это тензор формы H x W (семантическая) или H x W x N instances…
Текущий релиз AlexBred.com: первые 100 уроков русскоязычной программы.
Содержание урока
- Цели обучения
- Проблема
- Концепция
- Семантическая, экземплярная и паноптическая сегментация
- Форма U-Net
- Транспонированная свёртка и билинейное увеличение масштаба
- Cross-entropy на сетке пикселей
- Dice loss и почему она нужна
- Метрики оценки
- Компромисс входного разрешения
- Соберите сами
- Шаг 1: Блок кодировщика
- Шаг 2: Блоки понижения и повышения разрешения
- Шаг 3: U-Net
- Шаг 4: Функции потерь
- Шаг 5: Метрика IoU
- Шаг 6: Синтетический датасет для сквозной проверки
- Шаг 7: Цикл обучения
- Используйте
- Подготовьте к поставке
- Упражнения
- Ключевые термины
- Дополнительные материалы
Сегментация — это классификация для каждого пикселя. U-Net делает её рабочей, объединяя кодировщик с понижением разрешения и декодировщик с повышением разрешения, а также соединяя их skip connections.
Тип: Сборка Языки: Python Предварительные требования: Фаза 4, урок 03 (CNN), фаза 4, урок 04 (Классификация изображений) Время: ~75 минут
Цели обучения
- Различать семантическую, экземплярную и паноптическую сегментацию и выбирать подходящую задачу для конкретной проблемы
- Собрать U-Net с нуля в PyTorch: блоки кодировщика, bottleneck, декодировщик с транспонированными свёртками и skip connections
- Реализовать попиксельную cross-entropy, Dice loss и комбинированную функцию потерь — нынешний стандарт по умолчанию для медицинской и промышленной сегментации
- Читать метрики IoU и Dice по классам и определять, вызван ли плохой результат полнотой для малых объектов, точностью границ или дисбалансом классов
Проблема
Классификация выдаёт одну метку на изображение. Детекция выдаёт несколько рамок на изображение. Сегментация выдаёт одну метку на пиксель. Для входа размера H x W выход — это тензор формы H x W (семантическая) или H x W x N_instances (экземплярная). Это миллионы предсказаний на изображение, а не одно.
Именно структура сегментации делает её основой почти любого продукта компьютерного зрения с плотным предсказанием: медицинская визуализация (маски опухолей), автономное вождение (дорога, полоса, препятствие), спутниковые снимки (контуры зданий, границы полей), разбор документов (зоны макета), робототехника (области для захвата). Ни одну из этих задач нельзя решить рамкой вокруг объекта: нужен точный силуэт.
Архитектурную проблему легко сформулировать, но сложно решить: сети одновременно нужны глобальный контекст изображения (какого типа это сцена) и локальная детализация пикселей (какой именно пиксель — дорога, а какой — тротуар). Обычная CNN сжимает пространственное представление, чтобы получить контекст, и отбрасывает детали. U-Net стала конструкцией, которая сохранила и то и другое.
Концепция
Семантическая, экземплярная и паноптическая сегментация
- Семантическая говорит: «этот пиксель — дорога, тот пиксель — автомобиль». Два стоящих рядом автомобиля сливаются в одну область.
- Экземплярная говорит: «этот пиксель — автомобиль №3, тот пиксель — автомобиль №5». Она игнорирует фоновую среду (“stuff” = небо, дорога, трава).
- Паноптическая объединяет обе: каждый пиксель получает метку класса, каждый экземпляр — уникальный id, сегментируются и среда, и объекты.
Этот урок посвящён семантической сегментации. Следующий урок (Mask R-CNN) посвящён экземплярной.
Форма U-Net
Кодировщик в четыре шага вдвое уменьшает пространственное разрешение и удваивает число каналов. Декодировщик делает обратное: в четыре шага удваивает пространственное разрешение и уменьшает число каналов вдвое. Skip connections конкатенируют совпадающие по разрешению признаки кодировщика с признаками декодировщика на каждом уровне. Последняя свёртка 1x1 отображает 64 -> num_classes при полном разрешении.
Почему skip connections необходимы: к тому моменту, когда декодировщик пытается выдать попиксельные предсказания, он видел лишь небольшие карты признаков. Без пропусков он не может точно локализовать границы, поскольку эта информация была сжата в кодировщике. Skip connections передают ему карты признаков высокого разрешения, вычисленные кодировщиком на пути вниз.
Транспонированная свёртка и билинейное увеличение масштаба
Декодировщик должен расширять пространственные размеры. Есть два варианта:
- Транспонированная свёртка (
nn.ConvTranspose2d) — обучаемое увеличение масштаба. Исторический вариант U-Net по умолчанию. Может создавать артефакты в виде шахматной доски, если stride и размер ядра не делятся ровно. - Билинейное увеличение масштаба + свёртка 3x3 — плавное увеличение масштаба с последующей свёрткой. Меньше артефактов, меньше параметров; современный вариант по умолчанию.
Оба варианта встречаются на практике. Для первой U-Net безопаснее билинейный.
Cross-entropy на сетке пикселей
Для семантической сегментации с C классами выход модели имеет форму (N, C, H, W). Цель имеет форму (N, H, W) с целочисленными id классов. Cross-entropy идентична случаю классификации, но применяется в каждой пространственной позиции:
Loss = mean over (n, h, w) of -log( softmax(logits[n, :, h, w])[target[n, h, w]] )
F.cross_entropy в PyTorch нативно обрабатывает эту форму. Изменять форму тензоров не нужно.
Dice loss и почему она нужна
Cross-entropy рассматривает каждый пиксель одинаково. Это неверно, когда один класс доминирует в кадре (медицинская визуализация: 99% фона, 1% опухоли). Сеть может получить точность 99%, предсказывая фон везде, и при этом быть бесполезной.
Dice loss решает это, напрямую оптимизируя перекрытие предсказанной и истинной масок:
Dice(p, y) = 2 * sum(p * y) / (sum(p) + sum(y) + epsilon)
Dice_loss = 1 - Dice
где p — карта вероятностей sigmoid/softmax для класса, а y — бинарная ground-truth маска. Потеря равна нулю только при идеальном перекрытии. Поскольку она основана на отношении, дисбаланс классов несущественен.
На практике используйте комбинированную функцию потерь:
L = L_cross_entropy + lambda * L_dice (lambda ~ 1)
Cross-entropy даёт стабильные градиенты в начале обучения; Dice направляет завершающую часть обучения на реальное совпадение формы маски. Эта комбинация — стандарт медицинской визуализации, который трудно превзойти на любом датасете с дисбалансом классов.
Метрики оценки
- Попиксельная точность — процент корректно предсказанных пикселей. Дёшево. На несбалансированных данных не работает по той же причине, что и accuracy в классификации.
- IoU по классам — intersection over union для маски каждого класса; среднее по классам = mIoU.
- Dice (F1 по пикселям) — похожа на IoU;
Dice = 2 * IoU / (1 + IoU). В медицинской визуализации предпочитают Dice, в сообществе автономного вождения — IoU; они монотонно связаны. - Boundary F1 — измеряет близость предсказанных границ к ground-truth границам, штрафуя даже малые смещения. Важна для высокоточных задач, например инспекции полупроводников.
Сообщайте IoU по классам, а не только mIoU. Средняя IoU скрывает класс с 15%, когда девять других имеют 85%.
Компромисс входного разрешения
Кодировщик U-Net в четыре раза делит разрешение пополам, поэтому вход должен делиться на 16. Медицинские изображения часто имеют размер 512x512 или 1024x1024. Обрезки для автономного вождения — 2048x1024. Затраты памяти U-Net масштабируются как H * W * C_max, и при 1024x1024 с 1024 каналами bottleneck один прямой проход уже использует гигабайты VRAM.
Два стандартных обходных решения:
- Разбить вход на тайлы — обработать тайлы 256x256 с перекрытием и сшить результат.
- Заменить bottleneck дилатированными свёртками, которые сохраняют большее пространственное разрешение, но расширяют receptive field (семейство DeepLab).
Для первой модели U-Net с входом 256x256 и базой в 64 канала комфортно обучается на 8 GB VRAM.
Соберите сами
Шаг 1: Блок кодировщика
Две свёртки 3x3 с batch norm и ReLU. Первая свёртка меняет число каналов; вторая сохраняет его.
import torch
import torch.nn as nn
import torch.nn.functional as F
class DoubleConv(nn.Module):
def __init__(self, in_c, out_c):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(in_c, out_c, kernel_size=3, padding=1, bias=False),
nn.BatchNorm2d(out_c),
nn.ReLU(inplace=True),
nn.Conv2d(out_c, out_c, kernel_size=3, padding=1, bias=False),
nn.BatchNorm2d(out_c),
nn.ReLU(inplace=True),
)
def forward(self, x):
return self.net(x)
Этот блок используется во всей сети. bias=False, потому что beta в BN берёт на себя роль смещения.
Шаг 2: Блоки понижения и повышения разрешения
class Down(nn.Module):
def __init__(self, in_c, out_c):
super().__init__()
self.net = nn.Sequential(
nn.MaxPool2d(2),
DoubleConv(in_c, out_c),
)
def forward(self, x):
return self.net(x)
class Up(nn.Module):
def __init__(self, in_c, out_c):
super().__init__()
self.up = nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False)
self.conv = DoubleConv(in_c, out_c)
def forward(self, x, skip):
x = self.up(x)
if x.shape[-2:] != skip.shape[-2:]:
x = F.interpolate(x, size=skip.shape[-2:], mode="bilinear", align_corners=False)
x = torch.cat([skip, x], dim=1)
return self.conv(x)
Проверка только пространственной формы (shape[-2:]) обрабатывает входы, чьи размеры не делятся на 16; безопасный F.interpolate выравнивает тензор перед конкатенацией. Сравнение полной формы сработало бы и при различии числа каналов — это должно быть заметной ошибкой, а не тихим интерполированием.
Шаг 3: U-Net
class UNet(nn.Module):
def __init__(self, in_channels=3, num_classes=2, base=64):
super().__init__()
self.inc = DoubleConv(in_channels, base)
self.d1 = Down(base, base * 2)
self.d2 = Down(base * 2, base * 4)
self.d3 = Down(base * 4, base * 8)
self.d4 = Down(base * 8, base * 16)
self.u1 = Up(base * 16 + base * 8, base * 8)
self.u2 = Up(base * 8 + base * 4, base * 4)
self.u3 = Up(base * 4 + base * 2, base * 2)
self.u4 = Up(base * 2 + base, base)
self.outc = nn.Conv2d(base, num_classes, kernel_size=1)
def forward(self, x):
x1 = self.inc(x)
x2 = self.d1(x1)
x3 = self.d2(x2)
x4 = self.d3(x3)
x5 = self.d4(x4)
x = self.u1(x5, x4)
x = self.u2(x, x3)
x = self.u3(x, x2)
x = self.u4(x, x1)
return self.outc(x)
net = UNet(in_channels=3, num_classes=2, base=32)
x = torch.randn(1, 3, 256, 256)
print(f"output: {net(x).shape}")
print(f"params: {sum(p.numel() for p in net.parameters()):,}")
Форма выхода (1, 2, 256, 256) — то же пространственное разрешение, что у входа, и num_classes каналов. При base=32 — примерно 7,7 млн параметров.
Шаг 4: Функции потерь
def dice_loss(logits, targets, num_classes, eps=1e-6):
probs = F.softmax(logits, dim=1)
targets_one_hot = F.one_hot(targets, num_classes).permute(0, 3, 1, 2).float()
dims = (0, 2, 3)
intersection = (probs * targets_one_hot).sum(dim=dims)
denom = probs.sum(dim=dims) + targets_one_hot.sum(dim=dims)
dice = (2 * intersection + eps) / (denom + eps)
return 1 - dice.mean()
def combined_loss(logits, targets, num_classes, lam=1.0):
ce = F.cross_entropy(logits, targets)
dc = dice_loss(logits, targets, num_classes)
return ce + lam * dc, {"ce": ce.item(), "dice": dc.item()}
Dice вычисляется для каждого класса, затем усредняется (macro Dice). eps предотвращает деление на ноль для классов, отсутствующих в батче.
Шаг 5: Метрика IoU
@torch.no_grad()
def iou_per_class(logits, targets, num_classes):
preds = logits.argmax(dim=1)
ious = torch.zeros(num_classes)
for c in range(num_classes):
pred_c = (preds == c)
true_c = (targets == c)
inter = (pred_c & true_c).sum().float()
union = (pred_c | true_c).sum().float()
ious[c] = (inter / union) if union > 0 else torch.tensor(float("nan"))
return ious
Возвращает вектор длины C. nan отмечает классы, отсутствующие в батче, — не усредняйте по ним при вычислении mIoU.
Шаг 6: Синтетический датасет для сквозной проверки
Генерируйте фигуры на цветных фонах, чтобы сети пришлось выучить форму, а не цвет пикселей.
import numpy as np
from torch.utils.data import Dataset, DataLoader
def synthetic_segmentation(num_samples=200, size=64, seed=0):
rng = np.random.default_rng(seed)
images = np.zeros((num_samples, size, size, 3), dtype=np.float32)
masks = np.zeros((num_samples, size, size), dtype=np.int64)
for i in range(num_samples):
bg = rng.uniform(0, 1, (3,))
images[i] = bg
masks[i] = 0
num_shapes = rng.integers(1, 4)
for _ in range(num_shapes):
cls = int(rng.integers(1, 3))
color = rng.uniform(0, 1, (3,))
cx, cy = rng.integers(10, size - 10, size=2)
r = int(rng.integers(4, 12))
yy, xx = np.meshgrid(np.arange(size), np.arange(size), indexing="ij")
if cls == 1:
mask = (xx - cx) ** 2 + (yy - cy) ** 2 < r ** 2
else:
mask = (np.abs(xx - cx) < r) & (np.abs(yy - cy) < r)
images[i][mask] = color
masks[i][mask] = cls
images[i] += rng.normal(0, 0.02, images[i].shape)
images[i] = np.clip(images[i], 0, 1)
return images, masks
class SegDataset(Dataset):
def __init__(self, images, masks):
self.images = images
self.masks = masks
def __len__(self):
return len(self.images)
def __getitem__(self, i):
img = torch.from_numpy(self.images[i]).permute(2, 0, 1).float()
mask = torch.from_numpy(self.masks[i]).long()
return img, mask
Три класса: фон (0), круги (1), квадраты (2). Сеть должна научиться различать форму.
Шаг 7: Цикл обучения
def train_one_epoch(model, loader, optimizer, device, num_classes):
model.train()
loss_sum, total = 0.0, 0
iou_sum = torch.zeros(num_classes)
for x, y in loader:
x, y = x.to(device), y.to(device)
logits = model(x)
loss, _ = combined_loss(logits, y, num_classes)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_sum += loss.item() * x.size(0)
total += x.size(0)
iou_sum += iou_per_class(logits, y, num_classes).nan_to_num(0)
return loss_sum / total, iou_sum / len(loader)
Запустите это на 10–30 эпох на синтетическом датасете и наблюдайте, как mIoU для классов фигур поднимается выше 0,9. Обратите внимание: nan_to_num(0) рассматривает классы, отсутствующие в батче, как ноль; для точной IoU по классам маскируйте по наличию и используйте torch.nanmean по батчам при оценке вместо усреднения здесь.
Используйте
Для промышленного применения segmentation_models_pytorch (“smp”) оборачивает каждую стандартную архитектуру сегментации с любым backbone из torchvision или timm. Три строки:
import segmentation_models_pytorch as smp
model = smp.Unet(
encoder_name="resnet34",
encoder_weights="imagenet",
in_channels=3,
classes=3,
)
Также полезно знать для реальной работы:
- DeepLabV3+ заменяет понижение разрешения на основе max-pool дилатированными свёртками, поэтому bottleneck сохраняет разрешение; обеспечивает более точные границы на спутниковых данных и вождении.
- SegFormer заменяет свёрточный кодировщик иерархическим трансформером; текущий SOTA на многих бенчмарках.
- Mask2Former / OneFormer объединяют семантическую, экземплярную и паноптическую сегментацию в единой архитектуре.
Все три являются заменами в smp или transformers с тем же data loader.
Подготовьте к поставке
Этот урок создаёт:
outputs/prompt-segmentation-task-picker.md— промпт, выбирающий семантическую, экземплярную или паноптическую сегментацию и называющий архитектуру для конкретной задачи.outputs/skill-segmentation-mask-inspector.md— навык, сообщающий распределение классов, статистику предсказанных масок и классы, которые недопредсказываются или имеют размытые границы.
Упражнения
- (Легко) Реализуйте
bce_dice_lossдля задачи бинарной сегментации (передний план против фона). Проверьте на синтетическом двухклассовом датасете, что комбинированная функция потерь сходится быстрее, чем один BCE, когда передний план составляет 5% пикселей. - (Средне) Замените up-block
nn.Upsample + convна up-blocknn.ConvTranspose2d. Обучите оба на синтетическом датасете и сравните mIoU. Наблюдайте, где в версии с транспонированной свёрткой появляются артефакты в виде шахматной доски. - (Сложно) Возьмите реальный датасет сегментации (Oxford-IIIT Pets, мини-раздел Cityscapes или медицинское подмножество) и обучите U-Net до отставания не более чем на 2 пункта IoU от эталона
smp.Unet. Сообщите IoU по классам и определите, каким классам добавление Dice в функцию потерь помогает сильнее всего.
Ключевые термины
| Термин | Как говорят | Что это на самом деле означает |
|---|---|---|
| Семантическая сегментация | «Разметить каждый пиксель» | Попиксельная классификация по C классам; экземпляры одного класса сливаются |
| Экземплярная сегментация | «Разметить каждый объект» | Разделяет разные экземпляры одного класса; только передний план |
| Паноптическая сегментация | «Семантика + экземпляры» | Каждый пиксель получает класс; каждый экземпляр объекта также получает уникальный id |
| Skip connection | «Мост U-Net» | Конкатенация признаков кодировщика с признаками декодировщика совпадающего разрешения; сохраняет высокочастотные детали |
| Транспонированная свёртка | «Деконволюция» | Обучаемое увеличение масштаба; может создавать артефакты в виде шахматной доски |
| Dice loss | «Потеря перекрытия» | 1 - 2 |
| mIoU | «Среднее intersection over union» | Средняя IoU по классам; стандартная метрика сегментации в сообществе |
| Boundary F1 | «Точность границ» | F1, вычисленная только по граничным пикселям; важна для задач, требующих высокой точности |
Дополнительные материалы
- U-Net: Convolutional Networks for Biomedical Image Segmentation (Ronneberger et al., 2015) — оригинальная статья; фигура, которую все копируют, находится на странице 2
- Fully Convolutional Networks (Long et al., 2015) — статья, впервые сделавшая сегментацию сквозной свёрточной задачей
- segmentation_models_pytorch — эталон для производственной сегментации: все стандартные архитектуры и функции потерь
- Lessons learned from training SOTA segmentation (kaggle.com competitions) — разбор того, почему TTA, псевдометки и веса классов важны на реальных данных
Источник: Semantic Segmentation — U-Net 04.06 — Обнаружение объектов — YOLO с нуля · Фаза 04 — Компьютерное зрение · 04.08 — Экземплярная сегментация — Mask R-CNN · Полный каталог