U-Net 语义分割 本节摘要:分割,就是对每个像素做分类。U-Net 让它变得可行:把一个下采样的编码器与一个上采样的解码器配对,在两者之间用跳跃连接把高分辨率细节送回去。本节从零在 PyTorch 里搭一个 U-Net——编码块、瓶颈、用转置卷积(或双线性上采样)的解码器、skip 连接,讲清为什么 skip 是必要的。我们实现像素级交叉熵、Dice 损失以及「CE + λ·Dice」的组合损失(医学与工业分割的当前默认),并读 IoU/Dice 每类指标,诊断一个低分究竟来自小目标召回、边界精度还是类别不平衡。读完本节,你能区分语义/实例/全景分割,并为给定问题挑对任务。 对应原课程:Phase 4 · Lesson 07 · (原英文 )。
本节摘要:分割,就是对每个像素做分类。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)。
阅读完本节,你应当能够:
分类每图输出一个标签,检测每图输出几个框,分割每个像素输出一个标签。对 H × W 的输入,输出是 H × W(语义)或 H × W × N_instances(实例)的张量。每图数百万次预测,而不是一次。
分割的结构正是它驱动几乎所有密集预测视觉产品的原因:医学影像(肿瘤掩码)、自动驾驶(路面、车道、障碍)、卫星(建筑轮廓、作物边界)、文档解析(版面区域)、机器人(可抓取区域)。这些任务没有一个能用「圈个框」解决,它们需要精确的轮廓。
架构问题说起来简单,解决不简单:你需要网络同时看到全局上下文(这是什么场景)和局部像素细节(到底哪个像素是路、哪个是人行道)。标准 CNN 为换上下文而压缩空间,从而丢掉细节。U-Net 是同时拿到两者的设计。
本节讲语义。下一节(Mask R-CNN)讲实例。
编码器四次减半空间分辨率、加倍通道。解码器反过来:四次加倍分辨率、减半通道。skip 连接在每个分辨率上把匹配的编码器特征与解码器特征拼接。最后的 1×1 卷积在全分辨率上把 64 -> num_classes 映射。
为什么 skip 必要:解码器在试图输出像素级预测时,只见过小特征图。没有 skip,它无法准确定位边缘,因为那信息在编码器里被压没了。skip 把编码器下行路上算出的高分辨率特征图递给它。
解码器要扩大空间维度。两个选择:
nn.ConvTranspose2d)—— 可学习的上采样。U-Net 的历史默认。步幅和核尺寸不能整除时会产生棋盘伪影。两者生产中都见。第一个 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。
交叉熵平等对待每个像素。当一类主导画面(医学影像: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 把训练后期聚焦到真正匹配掩码形状上。这个组合是医学影像的默认,在任何类别不平衡的数据集上都难被超越。
Dice = 2 * IoU / (1 + IoU)。医学偏好 Dice,自动驾驶偏好 IoU;两者单调相关。报告每类 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 显存。
两种标准折中:
第一个模型用 256×256 输入、64 通道基的 U-Net,在 8 GB 显存上舒适训练。
两个 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 已承担偏置。
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 在拼接前对齐张量。比完整形状也会因通道数差异触发,而那应该是个响亮的错误,不是静默插值。
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 万参数。
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 防止批次中缺席的类除零。
@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 时不要把它们算进去。
在彩色背景上生成形状,让网络必须学形状,而非像素颜色。
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)。网络必须学会区分形状。
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, )
真实工作中还值得知道:
这三者在 smp 或 transformers 里都是即插即用,用同一个 dataloader。
本节产出两个可复用文件(位于原课程 outputs/):
prompt-segmentation-task-picker.md:一个提示词——在语义/实例/全景分割之间挑,并为给定任务命名架构。skill-segmentation-mask-inspector.md:一个技能——报告类别分布、预测掩码统计、哪些类被低预测或边界模糊。bce_dice_loss。在合成二类数据集上验证:当前景占 5% 像素时,组合损失比纯 BCE 收敛更快。nn.Upsample + conv 上采样块换成 nn.ConvTranspose2d 上采样块。在合成数据集上训两者,比较 mIoU。观察转置卷积版在哪里出现棋盘伪影。smp.Unet 参考差 2 个 IoU 点以内。报告每类 IoU,找出哪些类从加 Dice 损失中获益最多。H×W(语义)或 H×W×N_instances(实例),每图数百万预测。1 - 2|A∩B|/(|A|+|B|),基于比值,对类别不平衡鲁棒。下一节,我们离开「同类像素合并」,进入实例分割——用 Mask R-CNN 讲清如何为每个物体单独预测一个像素掩码。