U-Net 语义分割


文档摘要

U-Net 语义分割 本节摘要:分割,就是对每个像素做分类。U-Net 让它变得可行:把一个下采样的编码器与一个上采样的解码器配对,在两者之间用跳跃连接把高分辨率细节送回去。本节从零在 PyTorch 里搭一个 U-Net——编码块、瓶颈、用转置卷积(或双线性上采样)的解码器、skip 连接,讲清为什么 skip 是必要的。我们实现像素级交叉熵、Dice 损失以及「CE + λ·Dice」的组合损失(医学与工业分割的当前默认),并读 IoU/Dice 每类指标,诊断一个低分究竟来自小目标召回、边界精度还是类别不平衡。读完本节,你能区分语义/实例/全景分割,并为给定问题挑对任务。 对应原课程:Phase 4 · Lesson 07 · (原英文 )。

U-Net 语义分割

本节摘要:分割,就是对每个像素做分类。U-Net 让它变得可行:把一个下采样的编码器与一个上采样的解码器配对,在两者之间用跳跃连接把高分辨率细节送回去。本节从零在 PyTorch 里搭一个 U-Net——编码块、瓶颈、用转置卷积(或双线性上采样)的解码器、skip 连接,讲清为什么 skip 是必要的。我们实现像素级交叉熵、Dice 损失以及「CE + λ·Dice」的组合损失(医学与工业分割的当前默认),并读 IoU/Dice 每类指标,诊断一个低分究竟来自小目标召回、边界精度还是类别不平衡。读完本节,你能区分语义/实例/全景分割,并为给定问题挑对任务。

对应原课程:Phase 4 · Lesson 07 · semantic-segmentation-unet(原英文 phases/04-computer-vision/07-semantic-segmentation-unet/docs/en.md)。

学习目标

阅读完本节,你应当能够:

  1. 区分语义、实例、全景分割,并为给定问题挑对任务。
  2. 在 PyTorch 中从零搭一个 U-Net:编码块、瓶颈、用转置卷积的解码器、skip 连接。
  3. 实现像素级交叉熵、Dice 损失、组合损失(医学与工业分割的当前默认)。
  4. 读每类的 IoU 与 Dice 指标,诊断一个低分究竟来自小目标召回、边界精度还是类别不平衡。

一、问题与直觉

分类每图输出一个标签,检测每图输出几个框,分割每个像素输出一个标签。对 H × W 的输入,输出是 H × W(语义)或 H × W × N_instances(实例)的张量。每图数百万次预测,而不是一次。

分割的结构正是它驱动几乎所有密集预测视觉产品的原因:医学影像(肿瘤掩码)、自动驾驶(路面、车道、障碍)、卫星(建筑轮廓、作物边界)、文档解析(版面区域)、机器人(可抓取区域)。这些任务没有一个能用「圈个框」解决,它们需要精确的轮廓

架构问题说起来简单,解决不简单:你需要网络同时看到全局上下文(这是什么场景)和局部像素细节(到底哪个像素是路、哪个是人行道)。标准 CNN 为换上下文而压缩空间,从而丢掉细节。U-Net 是同时拿到两者的设计。

语义 vs 实例 vs 全景

  • 语义 说「这个像素是路,那个像素是车」。两辆挨着的车合并成一团。
  • 实例 说「这个像素是 3 号车,那个是 5 号车」。忽略背景 stuff(stuff = 天空、路、草地)。
  • 全景 统一两者:每个像素有类别标签,每个实例有唯一 id,stuff 和 things 都分割。

本节讲语义。下一节(Mask R-CNN)讲实例。

U-Net 的形状

编码器四次减半空间分辨率、加倍通道。解码器反过来:四次加倍分辨率、减半通道。skip 连接在每个分辨率上把匹配的编码器特征与解码器特征拼接。最后的 1×1 卷积在全分辨率上把 64 -> num_classes 映射。

为什么 skip 必要:解码器在试图输出像素级预测时,只见过小特征图。没有 skip,它无法准确定位边缘,因为那信息在编码器里被压没了。skip 把编码器下行路上算出的高分辨率特征图递给它。

转置卷积 vs 双线性上采样

解码器要扩大空间维度。两个选择:

  • 转置卷积(nn.ConvTranspose2d)—— 可学习的上采样。U-Net 的历史默认。步幅和核尺寸不能整除时会产生棋盘伪影。
  • 双线性上采样 + 3×3 卷积 —— 平滑上采样后接一个卷积。伪影更少、参数更少,现代默认。

两者生产中都见。第一个 U-Net 用双线性更安全。

像素网格上的交叉熵

对 C 类的语义分割,模型输出 (N, C, H, W)。目标是 (N, H, W) 的整数类 ID。交叉熵与分类情况一致,只是作用在每个空间位置:

Loss = 对 (n, h, w) 求 -log( softmax(logits[n, :, h, w])[target[n, h, w]] ) 的均值

PyTorch 的 F.cross_entropy 原生处理这个形状,无需 reshape。

Dice 损失及其必要性

交叉熵平等对待每个像素。当一类主导画面(医学影像:99% 背景,1% 肿瘤)时这就错了——网络处处预测背景就能拿 99% 准确率,却毫无用处。

Dice 损失直接优化预测与真实掩码的重叠来解决:

Dice(p, y) = 2 * sum(p * y) / (sum(p) + sum(y) + epsilon) Dice_loss = 1 - Dice

其中 p 是某类的 sigmoid/softmax 概率图,y 是二值真实掩码。只有重叠完美时损失才为零。因为它基于比值,类别不平衡无关紧要。

实践里用组合损失:

L = L_cross_entropy + lambda * L_dice (lambda ~ 1)

交叉熵在训练早期给稳定梯度;Dice 把训练后期聚焦到真正匹配掩码形状上。这个组合是医学影像的默认,在任何类别不平衡的数据集上都难被超越。

评估指标

  • 像素准确率 —— 预测对的像素百分比。便宜。在不平衡数据上失效,原因与分类中的准确率一样。
  • 每类 IoU —— 每类掩码的交并比;跨类平均 = mIoU。
  • Dice(像素上的 F1) —— 与 IoU 类似;Dice = 2 * IoU / (1 + IoU)。医学偏好 Dice,自动驾驶偏好 IoU;两者单调相关。
  • 边界 F1 —— 衡量预测边界与真实边界有多近,即使小偏移也惩罚。对半导体检测这类高精度任务重要。

报告每类 IoU,不只是 mIoU。九类 85%、一类 15% 时,mIoU 会藏起那 15%。

输入分辨率权衡

U-Net 编码器四次减半分辨率,所以输入必须能被 16 整除。医学图像常 512×512 或 1024×1024。自动驾驶裁剪 2048×1024。U-Net 显存随 H * W * C_max 增长,1024×1024、1024 瓶颈通道时一次前向已用数 GB 显存。

两种标准折中:

  1. 分块(tile)—— 以带重叠的 256×256 块处理再拼接。
  2. 空洞卷积替换瓶颈,保持更高分辨率但扩大感受野(DeepLab 家族)。

第一个模型用 256×256 输入、64 通道基的 U-Net,在 8 GB 显存上舒适训练。

二、从零实现

步骤 1:编码块

两个 3×3 卷积加 BN 和 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 因为 BN 的 beta 已承担偏置。

步骤 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 时约 770 万参数。

步骤 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 个 epoch,看形状类的 mIoU 爬过 0.9。注意 nan_to_num(0) 把批次中缺席的类当零;要准确每类 IoU,评估时按是否存在掩码,跨批用 torch.nanmean 而非在这里平均。

三、框架对比

生产中,segmentation_models_pytorch("smp")把每个标准分割架构 + 任意 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 下采样,让瓶颈保留分辨率;在卫星与驾驶数据上边界更利。
  • SegFormer 用层级 transformer 换掉卷积编码器;在许多基准上是当前 SOTA。
  • Mask2Former / OneFormer 在单一架构里统一语义、实例、全景分割。

这三者在 smptransformers 里都是即插即用,用同一个 dataloader。

四、可复用产物

本节产出两个可复用文件(位于原课程 outputs/):

  • prompt-segmentation-task-picker.md:一个提示词——在语义/实例/全景分割之间挑,并为给定任务命名架构。
  • skill-segmentation-mask-inspector.md:一个技能——报告类别分布、预测掩码统计、哪些类被低预测或边界模糊。

五、练习

  1. (简单) 实现二分割任务(前景 vs 背景)的 bce_dice_loss。在合成二类数据集上验证:当前景占 5% 像素时,组合损失比纯 BCE 收敛更快。
  2. (中等)nn.Upsample + conv 上采样块换成 nn.ConvTranspose2d 上采样块。在合成数据集上训两者,比较 mIoU。观察转置卷积版在哪里出现棋盘伪影。
  3. (困难) 取一个真实分割数据集(Oxford-IIIT Pets、Cityscapes mini 子集,或一个医学子集),把 U-Net 训到与 smp.Unet 参考差 2 个 IoU 点以内。报告每类 IoU,找出哪些类从加 Dice 损失中获益最多。

本节要点回顾

  1. 分割是每像素分类——输出 H×W(语义)或 H×W×N_instances(实例),每图数百万预测。
  2. 语义/实例/全景三分——语义合并同类实例,实例只分前景物体,全景统一两者。
  3. U-Net = 编码器+解码器+skip——编码四次减半分辨率加倍通道,解码反过来,skip 在每个分辨率上拼接匹配特征。
  4. skip 是必要的——解码器只见过小特征图,没有 skip 就无法准确定位边缘。
  5. 上采样两选择——转置卷积(可学习,可能棋盘伪影)、双线性+卷积(平滑,现代默认)。
  6. 交叉熵平等对待像素——类别不平衡时失效;网络处处预测背景就能拿 99% 准确率。
  7. Dice 直接优化重叠——1 - 2|A∩B|/(|A|+|B|),基于比值,对类别不平衡鲁棒。
  8. 组合损失是默认——CE + λ·Dice,CE 给早期稳定梯度,Dice 聚焦后期掩码形状;医学与工业分割的标准。
  9. 报告每类 IoU——mIoU 会藏起一类 15%、九类 85% 的情况。
  10. 生产级选择——smp 包各种架构+骨干;DeepLabV3+、SegFormer、Mask2Former 是进阶选项。

下一节,我们离开「同类像素合并」,进入实例分割——用 Mask R-CNN 讲清如何为每个物体单独预测一个像素掩码。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U