简介:本资源面向医学图像分割方向的深度学习学习者与研究者,提供基于Transformer-Unet的Synapse腹部多器官8类分割完整实战项目,覆盖主动脉、胆囊、脾、左肾、右肾、肝、胰腺、胃等类别,适合具备一定PyTorch基础、希望掌握Transformer与U-Net结合方案的中高级读者。压缩包共2000个文件,约252.37MB,其中1280个png与697个jpg为数据集切片及可视化图像,18个py脚本涵盖训练、评估与推理流程,另有txt说明与readme文档辅助上手。项目采用AdamW优化器、余弦退火学习率衰减与交叉熵损失,train脚本输出loss、iou曲线、学习率衰减曲线、训练日志及最优与最终权重;evaluate脚本计算测试集iou、recall、precision与像素准确率;predice脚本生成gt及gt+image掩膜图像。代码注释详尽,README提供训练自有数据的傻瓜式指引。项目训练100个epoch后测试集像素准确率达0.99,mean iou为0.84,已有1228人学习下载,可作为医学分割论文复现与工程落地的参考方案。
1. 从一张腹部 CT 说起:Transformer-Unet 在 Synapse 多器官分割里到底解决了什么
腹部 CT 的多器官分割,是医学图像分割里最容易被低估的一类任务。肝脏、脾、左右肾、胰腺、胆囊、胃、主动脉这 8 个结构,在 CT 上灰度接近、边界模糊,胰腺和胆囊经常贴着肠道,脾和肝在部分层面几乎连成一片。传统 Unet 靠卷积堆叠感受野,局部纹理抓得准,但跨器官的全局位置关系——比如左肾永远在主动脉左侧、脾永远在胃的后外侧——它得靠足够深的网络和足够大的 batch 才能隐式学到。Transformer-Unet 这类结构把自注意力塞进编码器或瓶颈层,就是冲着这个全局依赖去的。
这份资源是一套完整的 Synapse 多器官分割实战包,包含代码、数据集组织方式、训练结果,8 类标签,AdamW 优化器配余弦退火,交叉熵损失,训练 100 个 epoch,测试集 pixel accuracy 0.99、mean IoU 0.84。适合两类人:一类是想跑通一个医学分割 baseline 的算法工程师,另一类是手里有 B 超、CT 数据、想照着改自己数据集的研究生。下面按「结构怎么搭 → 数据怎么喂 → 训练怎么跑 → 指标怎么读 → 坑在哪」的顺序拆开讲。
2. Transformer-Unet 的结构选型:注意力加在编码器还是瓶颈层
2.1 为什么不是纯 Transformer,也不是纯 Unet
纯 Transformer 分割(比如 SETR 那一路)把图像切成 patch 序列,全局建模能力强,但医学图像标注量小、分辨率高,patch 化之后浅层细节丢得厉害,小器官像胆囊、胰腺很容易被整块漏掉。纯 Unet 反过来,细节保得住,但长程依赖要靠堆深度,训练不稳定。
Transformer-Unet 的折中思路是:卷积主干负责浅层纹理和边缘,Transformer 模块只放在编码器末端或瓶颈层,用自注意力替换掉原本的几层卷积。这样既保留了 Unet 的 skip connection 把浅层特征送到解码器,又在最抽象的特征层上做全局关系建模。常见做法是编码器前几级用残差卷积块,最后一级或瓶颈换成多头自注意力加 FFN,解码器仍然用转置卷积或双线性上采样逐级恢复。
选型上要盯住一个参数:注意力放在哪一级。放在瓶颈层,显存开销最小,8 类分割够用;放在编码器每一级,参数量和显存翻倍,小数据集上容易过拟合。这份资源走的是瓶颈层方案,对单卡 12G 左右的显存比较友好。
2.2 编码器、瓶颈、解码器的具体搭法
下面这段是编码器加瓶颈的核心结构,卷积块负责下采样,瓶颈层插入多头注意力。代码做了注释,可以直接对照自己的实现改。
import torch import torch.nn as nn class ConvBlock(nn.Module): """标准双层卷积:Conv-BN-ReLU 重复两次,用于编码器各级""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x) class BottleneckTransformer(nn.Module): """瓶颈层自注意力:把特征图展平成序列,做多头注意力再还原""" def __init__(self, dim, num_heads=8, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(dim) self.attn = nn.MultiheadAttention(dim, num_heads, dropout=dropout, batch_first=True) self.norm2 = nn.LayerNorm(dim) self.ffn = nn.Sequential( nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim), ) def forward(self, x): # x: [B, C, H, W] -> [B, H*W, C] B, C, H, W = x.shape seq = x.flatten(2).transpose(1, 2) seq = seq + self.attn(self.norm1(seq), self.norm1(seq), self.norm1(seq))[0] seq = seq + self.ffn(self.norm2(seq)) return seq.transpose(1, 2).view(B, C, H, W)逻辑说明:ConvBlock是编码器每一级的基本单元,两次 3x3 卷积把通道数翻上去、空间尺寸靠后面的池化降下来。BottleneckTransformer先把[B, C, H, W]展平成[B, H*W, C]的序列,做一次多头自注意力,再过 FFN,最后还原回特征图。num_heads=8是常见起点,通道数 512 时每个头 64 维;dropout=0.1在医学小数据集上建议保留,防止注意力权重过拟合到少数样本。
参数上要注意:瓶颈层特征图尺寸不能太大,否则H*W序列长度爆炸,注意力矩阵是平方复杂度。一般瓶颈层空间尺寸控制在 16x16 或 8x8,再大就得考虑窗口注意力或者下采样后再做。
2.3 解码器与 skip connection 的通道对齐
解码器每一级做两件事:上采样,然后和编码器对应层的特征拼接。拼接前通道数要对齐,常见做法是上采样后用 1x1 卷积把通道压到和 skip 特征一致,再 concat,再走一个ConvBlock。
class DecoderBlock(nn.Module): """解码器一级:上采样 -> 通道对齐 -> 拼接 skip -> 卷积融合""" def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up = nn.ConvTranspose2d(in_ch, out_ch, 2, stride=2) self.align = nn.Conv2d(skip_ch, out_ch, 1) self.fuse = ConvBlock(out_ch * 2, out_ch) def forward(self, x, skip): x = self.up(x) skip = self.align(skip) # 尺寸兜底:上采样后和 skip 差一个像素时做 padding if x.shape[-2:] != skip.shape[-2:]: x = nn.functional.interpolate(x, size=skip.shape[-2:], mode="bilinear", align_corners=False) return self.fuse(torch.cat([x, skip], dim=1))逻辑说明:ConvTranspose2d做 2 倍上采样,align用 1x1 卷积把 skip 的通道数压到和上采样结果一致,避免拼接后通道数失控。尺寸兜底那几行是血泪经验——输入尺寸不是 16 的整数倍时,下采样再上采样会对不齐,直接 concat 会报维度错误,用interpolate对齐最稳。
3. Synapse 数据集组织与训练脚本:从切片到 loss 曲线
3.1 数据目录结构与标签映射
Synapse 多器官分割的 8 类标签,常见映射是:0 背景、1 主动脉、2 胆囊、3 脾、4 左肾、5 右肾、6 肝、7 胰腺、8 胃。注意有的版本把脾和胃的顺序写反,训练前一定核对标签文件里的像素值分布,否则 IoU 会莫名其妙低一截。
目录组织建议按下面这样分,训练脚本按 split 读:
data/ ├── train/ │ ├── images/ # caseXXXX_sliceYYY.jpg │ └── masks/ # 同名 png,像素值为 0-8 ├── val/ │ ├── images/ │ └── masks/ └── test/ ├── images/ └── masks/项目正文里列出的case0033_slice070.jpg、case0010_slice063.jpg这些就是切片命名,case编号对应不同病人,slice是层号。同一个 case 的切片必须整体分到同一个 split,不能随机打散,否则相邻层几乎一样,验证集指标会虚高。这是医学分割里最常见的翻车点之一。
3.2 训练脚本的关键参数与 loss 曲线
训练脚本跑起来会输出训练集/验证集的 loss、IoU 曲线、学习率衰减曲线、训练日志和数据集可视化图像,最后保存最好和最后的权重。核心训练循环如下:
import torch from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR # 优化器:AdamW,weight_decay 是解耦权重衰减,比 Adam 的 L2 更稳 optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) # 余弦退火:从 1e-4 平滑降到 eta_min,100 epoch 对应 T_max scheduler = CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-6) # 损失:交叉熵,ignore_index=255 跳过未标注像素 criterion = torch.nn.CrossEntropyLoss(ignore_index=255) for epoch in range(100): model.train() for img, mask in train_loader: img, mask = img.cuda(), mask.cuda() pred = model(img) # [B, 9, H, W] loss = criterion(pred, mask) # mask: [B, H, W],值 0-8 optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() # 每个 epoch 后跑验证,记录 val loss 和 mean IoU,保存 best 权重逻辑说明:AdamW的weight_decay=1e-4是解耦衰减,和 Adam 里直接加 L2 不一样,医学分割上通常更稳。CosineAnnealingLR的T_max要等于总 epoch 数,eta_min=1e-6是学习率下限,别设成 0,否则后期几乎不更新。CrossEntropyLoss的ignore_index=255用来跳过没标注的像素,如果你的 mask 里没有 255,可以去掉这个参数。
参数怎么改:显存不够就把 batch size 降到 4 或 2,同时把学习率按比例降一点;训练不收敛先看学习率是不是太大,1e-4 对 AdamW 是常见起点,1e-3 容易震荡。验证 IoU 曲线如果一直低于训练 IoU 很多,优先怀疑数据划分泄漏,而不是模型容量。
3.3 评估脚本:IoU、Recall、Precision、像素准确率怎么算
评估脚本在测试集上算 mean IoU、recall、precision、pixel accuracy。多类分割的 IoU 是逐类算再平均,别用整体像素混淆矩阵直接除,否则背景类会拉高指标。
import numpy as np def compute_metrics(pred, gt, num_classes=9): """pred/gt: [H, W] 整数标签图,逐类算 IoU/Recall/Precision""" ious, recalls, precisions = [], [], [] for c in range(num_classes): p = (pred == c) g = (gt == c) inter = np.logical_and(p, g).sum() union = np.logical_or(p, g).sum() ious.append(inter / union if union > 0 else np.nan) recalls.append(inter / g.sum() if g.sum() > 0 else np.nan) precisions.append(inter / p.sum() if p.sum() > 0 else np.nan) return np.nanmean(ious), np.nanmean(recalls), np.nanmean(precisions)逻辑说明:逐类算完用np.nanmean平均,nan处理的是测试集里没出现的类,直接算 0 会拉低均值。背景类要不要算进 mean IoU,看你和谁比——和原论文比就按原论文的口径,自己内部对比就固定一种,别来回换。
3.4 推理脚本:生成 GT 与叠加掩膜
推理脚本对单张图输出预测掩膜,以及 GT 和 GT+image 的叠加图。叠加图用固定调色板,每个器官一个颜色,方便肉眼核对哪个类错得离谱。
import cv2 import numpy as np # 9 类调色板:背景黑,其余 8 类各一色 PALETTE = np.array([ [0, 0, 0], [255, 0, 0], [0, 255, 0], [0, 0, 255], [255, 255, 0], [255, 0, 255], [0, 255, 255], [128, 0, 0], [0, 128, 0], ], dtype=np.uint8) def overlay(image, mask, alpha=0.5): """image: BGR 原图,mask: [H, W] 标签图,返回叠加图""" color = PALETTE[mask] return cv2.addWeighted(image, 1 - alpha, color, alpha, 0)逻辑说明:PALETTE索引就是类别 id,mask直接当索引取色。alpha=0.5是叠加透明度,想看清边界可以调到 0.4。推理时记得把输入归一化和训练时保持一致,均值和方差对不上,预测会整体偏移。
4. 训练与推理的避坑排查:8 类分割最容易翻车的五个点
4.1 现象:验证 IoU 高得离谱,测试集一跑就崩
原因:同一个 case 的相邻切片被随机分到了训练集和验证集,相邻层几乎一样,验证集等于变相泄漏。解决:按 case 编号整体划分 split,训练/验证/测试三份的 case 不重叠,划分完打印一下三份的 case 列表核对。
4.2 现象:loss 一直不降,或者降到某个值就震荡
原因:学习率太大,或者CosineAnnealingLR的T_max设成了实际 epoch 的好几倍,学习率还没降下来训练就结束了。解决:先把 lr 降到 1e-4 甚至 5e-5 试一个 epoch,看 loss 是否稳定下降;T_max必须等于总 epoch 数,eta_min别设 0。
4.3 现象:小器官(胆囊、胰腺)IoU 接近 0,大器官正常
原因:交叉熵对类别不平衡不敏感,背景和大器官像素占绝大多数,小器官梯度被淹没。解决:常见做法是加 Dice loss 或 Focal loss 做加权,或者对小器官类别在交叉熵里设 class weight。这份资源用的是纯交叉熵,想提升小器官指标可以自己加一项 Dice。
4.4 现象:显存爆了,batch size 降到 1 还是 OOM
原因:瓶颈层注意力序列太长,H*W太大导致注意力矩阵平方级增长;或者输入分辨率没降,直接喂原图。解决:把瓶颈层之前的特征图尺寸控制住,输入统一 resize 到 224 或 256;实在不够就换窗口注意力,或者把注意力只放在更低分辨率的一级。
4.5 现象:推理叠加图整体偏移,预测掩膜和器官对不上
原因:推理时的归一化参数和训练时不一致,或者 mask 的标签映射和训练时反了(脾和胃顺序颠倒)。解决:把训练时的 mean/std 存进配置文件,推理脚本读同一份;标签映射写死在数据集类里,训练和推理共用,别两处各写一遍。
5. 把 0.84 mean IoU 再往上推:几个我实际会试的进阶手法
测试集 mean IoU 0.84、pixel accuracy 0.99,这个 baseline 已经能用了,但小器官还有空间。我一般会按下面顺序试,成本从低到高。
第一,损失函数加 Dice。交叉熵管像素分类,Dice 管区域重叠,两者按 0.5:0.5 加权,对小器官提升最直接。改法就是在训练循环里多算一项:
def dice_loss(pred, target, num_classes=9, eps=1e-6): """pred: [B, C, H, W] logits,target: [B, H, W]""" prob = torch.softmax(pred, dim=1) target_onehot = torch.nn.functional.one_hot(target, num_classes).permute(0, 3, 1, 2).float() dims = (0, 2, 3) inter = (prob * target_onehot).sum(dims) union = prob.sum(dims) + target_onehot.sum(dims) return 1 - ((2 * inter + eps) / (union + eps)).mean() # 组合损失 loss = 0.5 * criterion(pred, mask) + 0.5 * dice_loss(pred, mask)逻辑说明:dice_loss先把 logits 过 softmax 变成概率,target 做 one-hot,然后按通道和空间维度求和算 Dice 系数,eps防止除零。加权系数 0.5:0.5 是起点,小器官还是差就把 Dice 权重提到 0.7。
第二,数据增强。医学分割常用的有随机旋转、缩放、弹性形变、亮度对比度扰动。弹性形变对器官边界模拟最像,但别开太大,否则解剖结构变形过头反而有害。我一般旋转 ±15 度、缩放 0.9~1.1、亮度 ±0.1,弹性形变只在小器官上轻度用。
第三,深监督。在解码器每一级上采样后都接一个辅助分类头,算辅助 loss,加权求和。这样浅层也能拿到梯度,对小器官边界有帮助。辅助 loss 权重从 0.4 开始,训练后期可以降。
第四,测试时增强(TTA)。推理时对同一张图做水平翻转、多尺度缩放,预测结果平均。这个不改训练,只改推理,涨点稳定但推理时间翻几倍。对 0.84 这个量级,TTA 通常能再拿 0.5~1 个点。
验证方法上,我习惯固定一个测试集,每次改动只动一个变量,记录 mean IoU 和小器官逐类 IoU。别一次改三处,涨了不知道是哪处起作用,跌了也不知道该回退哪个。从那以后我每次调分割模型,都强制先跑一遍逐类 IoU 再决定下一步动哪里,不然盯着一个 mean 值来回试纯属浪费时间。希望帮到你。
本文还有配套的精品资源,点击获取