Что такое архитектура U-Net?

«Что такое архитектура U-Net?» — вопрос из категории Компьютерное зрение, который задают на 26% собеседований Data Scientist / ML Инженер. Ниже — развёрнутый ответ с разбором ключевых моментов.

Ответ

U-Net — это архитектура сверточной нейронной сети (CNN), изначально разработанная для биомедицинской семантической сегментации изображений. Ее название и характерная форма происходят от U-образной симметричной структуры, состоящей из пути сжатия (энкодера) и пути расширения (декодера).

Ключевые компоненты архитектуры:

  1. Энкодер (левая часть «U»):

    • Состоит из повторяющихся блоков, каждый из которых содержит две свертки 3x3 с активацией ReLU и последующей операцией макс-пулинга 2x2.
    • Цель: Извлечь контекстные, высокоуровневые признаки из изображения, постепенно уменьшая его пространственные размеры (ширину и высоту) и увеличивая глубину (количество карт признаков).
  2. Декодер (правая часть «U»):

    • Состоит из симметричных блоков, каждый из которых начинается с операции транспонированной свертки (transposed convolution) или upsampling для увеличения пространственных размеров.
    • За этим следуют две свертки 3x3 с ReLU.
    • Цель: Восстановить пространственную информацию для точного позиционирования границ объектов.
  3. Skip-connections (пропускные соединения):

    • Самая важная особенность U-Net.
    • Карты признаков с каждого уровня энкодера копируются и конкатенируются с соответствующими картами признаков на уровне декодера.
    • Зачем это нужно: При пулинге в энкодере теряется детальная пространственная информация (где именно находится объект). Skip-connections передают эти высокочастотные детали напрямую в декодер, что позволяет совместить «что» (контекст из энкодера) и «где» (детали из skip-connection) для точной локализации.

Упрощенная схема на PyTorch:

import torch
import torch.nn as nn
import torch.nn.functional as F

class DoubleConv(nn.Module):
    """Блок из двух сверток 3x3 с ReLU."""
    def __init__(self, in_channels, out_channels):
        super().__init__()
        self.double_conv = nn.Sequential(
            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_channels),
            nn.ReLU(inplace=True),
            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_channels),
            nn.ReLU(inplace=True)
        )
    def forward(self, x):
        return self.double_conv(x)

class UNet(nn.Module):
    def __init__(self, n_channels, n_classes):
        super().__init__()
        # Энкодер
        self.enc1 = DoubleConv(n_channels, 64)
        self.enc2 = DoubleConv(64, 128)
        self.pool = nn.MaxPool2d(2)
        # ... и т.д.
        # Декодер
        self.upconv4 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)
        self.dec4 = DoubleConv(512, 256) # 256 из upconv + 256 из skip = 512 вход
        # ... и т.д.
        self.out_conv = nn.Conv2d(64, n_classes, kernel_size=1)

    def forward(self, x):
        # Прямой проход энкодера с сохранением признаков для skip-connections
        e1 = self.enc1(x)
        e2 = self.enc2(self.pool(e1))
        e3 = self.enc3(self.pool(e2))
        # ...
        # Прямой проход декодера с конкатенацией
        d4 = self.upconv4(bottleneck)
        d4 = torch.cat([d4, e3], dim=1) # SKIP-CONNECTION здесь!
        d4 = self.dec4(d4)
        # ...
        return self.out_conv(d1)

Области применения: Помимо медицинской визуализации (сегментация клеток, органов на МРТ/КТ), U-Net успешно применяется для сегментации в автономном вождении, спутниковых снимках, восстановления изображений и colorization.