Transformer-Unet 实战 Synapse 多器官分割:8 类标签下的调参与避坑指南

发布时间:2026/10/2 17:57:46
Transformer-Unet 实战 Synapse 多器官分割:8 类标签下的调参与避坑指南 简介本资源面向医学图像分割方向的深度学习学习者与研究者提供基于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及gtimage掩膜图像。代码注释详尽README指导自定义数据训练已训练100个epoch测试集像素准确率0.99、mean iou 0.84已有1228人学习。1. Transformer-Unet 做 Synapse 多器官分割8 类标签下这套组合到底值不值得跑Synapse 多器官分割这个数据集做医学图像的人多少都听过。它来自腹部 CT标注了 8 个类别主动脉、胆囊、脾脏、左右肾、肝脏、胰腺、胃加上背景一共 9 类。难点不在类别多而在于器官之间的边界模糊、体积差异极大——肝脏能占几百个体素胆囊可能只有几十个而且不同患者之间的解剖结构差异明显。传统 Unet 在这个数据集上能跑出 0.75 左右的 Dice但小器官经常被吞掉边界也糊。Transformer-Unet 的思路很直接用 Transformer 的全局注意力补上 Unet 在长距离依赖上的短板。卷积核的感受野有限堆叠再深也只是局部到全局的渐进过程而自注意力从第一层就能让每个像素看到整张图。对于脾脏和胃这种位置相对固定但形状变化大的器官全局上下文确实有用。但代价也明显显存占用高、训练慢、小数据集容易过拟合。这篇面向的是想在这个方向落地的人——不管你是刚跑通 Unet 想往上加 Transformer还是已经看过 MissFormer、Swin-UNet 这类工作想自己复现一版。我会把数据准备、模型搭建、训练调参、结果验证这条链路拆开讲代码能直接抄参数会说明为什么这么设坑会提前标出来。8 类分割不是跑通就行小器官的 Dice 能不能稳住才是分水岭。2. Synapse 数据准备与 8 类标签的预处理从原始 CT 到可训练张量2.1 数据来源与目录结构Synapse 多器官数据集常见的有两个版本一个是 MICCAI 2015 挑战赛的多图谱标注数据另一个是后来整理的 Synapse 多器官版本通常包含 30 例左右的腹部 CT 扫描每例有对应的标注文件。原始格式一般是 NIfTI.nii.gz标注是整数标签图0 是背景1 到 8 对应八个器官。我一般会先把目录整理成下面这样后面写 Dataset 类的时候不用再改路径逻辑data/ ├── train/ │ ├── case0001_img.nii.gz │ ├── case0001_lbl.nii.gz │ ├── case0002_img.nii.gz │ ├── case0002_lbl.nii.gz │ └── ... ├── val/ │ ├── case0025_img.nii.gz │ ├── case0025_lbl.nii.gz │ └── ... └── test/ ├── case0028_img.nii.gz ├── case0028_lbl.nii.gz └── ...训练集和验证集按 8:2 或 7:3 切测试集留 2 到 3 例。注意同一个患者的不同扫描不能同时出现在训练和验证里否则 Dice 会虚高这个坑后面还会提。2.2 CT 值窗宽窗位与归一化腹部 CT 的原始 HU 值范围大概在 -1000 到 3000 之间直接送进网络梯度会炸。常见做法是先做窗宽窗位截断把腹部器官的 HU 范围映射到 [0,1]。腹部软组织一般在 -100 到 200 HU 之间肝脏脾脏大概在 40 到 60 HU所以窗宽设 400、窗位设 40 是比较通用的选择。import numpy as np import nibabel as nib def load_and_preprocess(img_path, lbl_path, window_level40, window_width400): img nib.load(img_path).get_fdata() lbl nib.load(lbl_path).get_fdata() # 窗宽窗位截断 lower window_level - window_width // 2 upper window_level window_width // 2 img np.clip(img, lower, upper) # 归一化到 [0,1] img (img - lower) / (upper - lower) # 标签转成 0-8 的整数 lbl lbl.astype(np.int64) return img, lbl窗宽窗位不是随便设的。设太窄器官内部对比度不够模型分不清脾脏和肝脏设太宽背景噪声进来小器官的信噪比下降。我试过 window_width300 和 500300 在胆囊上 Dice 会掉 2 个点左右500 对肾脏边界有轻微帮助但整体提升不明显。400 是折中。2.3 切片筛选与 2D 数据组织Synapse 是 3D 数据但 Transformer-Unet 做 2D 切片训练更常见原因是 3D 自注意力的显存开销太大30 例数据也不够训 3D 模型。2D 切片训练的关键是筛掉无效切片——腹部 CT 上下两端有很多纯背景或只有少量组织的切片这些切片送进去只会让模型偏向背景类。def filter_slices(img_volume, lbl_volume, min_organ_pixels100): valid_slices [] for i in range(img_volume.shape[2]): lbl_slice lbl_volume[:, :, i] organ_pixels np.sum(lbl_slice 0) if organ_pixels min_organ_pixels: valid_slices.append(i) return valid_slicesmin_organ_pixels 设 100 是个经验值。设太低很多只有几像素胆囊的切片进来小器官样本被稀释设太高胰腺这种本来就小的器官会被大量过滤掉。我一般会统计一下每个器官在多少切片里出现确保过滤后每个器官至少还有几百个切片的正样本。2.4 数据增强不能乱用的几种操作医学图像增强和自然图像不一样有些操作会破坏解剖合理性。水平翻转可以用因为腹部器官左右对称性不强但翻转后仍然合理垂直翻转要慎用肝脏翻到下面就不对了旋转角度不要超过 15 度大角度旋转会让器官位置关系失真。import random import numpy as np from scipy.ndimage import rotate def augment_slice(img, lbl): # 水平翻转 if random.random() 0.5: img np.fliplr(img) lbl np.fliplr(lbl) # 小角度旋转 if random.random() 0.3: angle random.uniform(-15, 15) img rotate(img, angle, reshapeFalse, order1) lbl rotate(lbl, angle, reshapeFalse, order0) # 弹性形变可选但参数要保守 return img.copy(), lbl.copy()弹性形变对医学图像有用但 alpha 和 sigma 要调小否则器官边界会被扭曲成不真实的形状。我一般 alpha10、sigma3 起步再大就要看视觉效果了。另外注意标签旋转要用最近邻插值order0用双线性会把整数标签变成小数后面算 loss 会出问题。3. Transformer-Unet 模型搭建编码器、瓶颈层与解码器怎么接3.1 整体结构选型为什么不是纯 Transformer纯 Transformer 分割模型比如 SETR在 Synapse 上也能跑但需要在大规模数据上预训练30 例数据从头训基本没戏。Transformer-Unet 的混合结构更务实编码器用卷积做浅层特征提取瓶颈层和深层用 Transformer 块做全局建模解码器再用卷积逐步恢复分辨率。具体来说我一般用 4 层编码器前两层是纯卷积后两层在卷积后接 Transformer 块。瓶颈层放 2 到 4 个 Transformer 块。解码器用转置卷积加跳跃连接和 Unet 一致。这样既有卷积的局部归纳偏置又有注意力的全局视野参数量也比纯 Transformer 小很多。3.2 Transformer 块的关键参数头数、维度、位置编码Transformer 块的核心参数是嵌入维度、注意力头数、MLP 扩展比。在 Synapse 这个规模上嵌入维度 256 到 512 比较合适头数 4 到 8MLP 扩展比 4。头数太多会导致每个头的维度太小注意力分布过于分散头数太少又退化成单头注意力全局建模能力下降。位置编码在 2D 分割里常用可学习的位置嵌入而不是正弦编码。原因是医学图像里器官的绝对位置有一定规律比如肝脏总在右上腹可学习嵌入能让模型记住这种先验。但要注意位置嵌入的分辨率要和特征图匹配训练和测试的输入尺寸不一致时要做插值。import torch import torch.nn as nn class TransformerBlock(nn.Module): def __init__(self, dim, num_heads8, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(dropout) ) def forward(self, x): # x: (B, N, C) x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x这里用的是 Pre-LN 结构训练比 Post-LN 稳定尤其在小数据集上不容易梯度爆炸。dropout 设 0.1 是常规操作如果过拟合严重可以加到 0.2但再高会影响收敛速度。3.3 编码器与解码器的拼接方式编码器每层输出两个东西卷积特征和 Transformer 特征。我一般把 Transformer 块的输出和卷积输出相加而不是拼接这样通道数不用变解码器那边不用改。跳跃连接传的是相加后的特征。class EncoderBlock(nn.Module): def __init__(self, in_ch, out_ch, use_transformerFalse, num_heads8): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) self.use_transformer use_transformer if use_transformer: self.transformer TransformerBlock(out_ch, num_heads) self.norm nn.LayerNorm(out_ch) def forward(self, x): x self.conv(x) if self.use_transformer: B, C, H, W x.shape x_flat x.flatten(2).transpose(1, 2) # (B, H*W, C) x_flat self.transformer(x_flat) x x_flat.transpose(1, 2).view(B, C, H, W) return x注意 flatten 之后序列长度是 H*W如果特征图是 64x64序列长度就是 4096自注意力的计算量是 4096 的平方显存占用不小。所以 Transformer 块一般只放在 16x16 或 32x32 的特征图上再大就要用窗口注意力或者线性注意力变体。3.4 损失函数Dice CE 的组合与类别权重Synapse 的类别极度不平衡背景占 90% 以上胆囊、胰腺这种小器官可能只占 0.5%。只用交叉熵模型会倾向于预测背景小器官的召回率很低。常见做法是 Dice Loss 加 Cross EntropyDice 负责优化重叠度CE 负责稳定梯度。class DiceCELoss(nn.Module): def __init__(self, num_classes9, class_weightsNone): super().__init__() self.num_classes num_classes self.ce nn.CrossEntropyLoss(weightclass_weights) def dice_loss(self, pred, target): pred torch.softmax(pred, dim1) dice 0 for c in range(1, self.num_classes): # 跳过背景 pred_c pred[:, c] target_c (target c).float() intersection (pred_c * target_c).sum() union pred_c.sum() target_c.sum() dice 1 - (2 * intersection 1e-5) / (union 1e-5) return dice / (self.num_classes - 1) def forward(self, pred, target): ce_loss self.ce(pred, target) dice_loss self.dice_loss(pred, target) return ce_loss dice_lossclass_weights 可以按类别频率的倒数来设但不要设得太极端否则背景的梯度会被压得太低模型在背景上反而出错。我一般把背景权重设 0.5小器官设 2 到 3中间器官设 1。这个权重需要根据验证集上的 per-class Dice 来调不是一次就能定下来的。4. 训练配置与调参学习率、批次大小、迭代次数怎么定4.1 优化器与学习率调度Transformer 类模型对优化器比较敏感AdamW 比 Adam 更稳因为权重衰减是解耦的。学习率初始值我一般设 1e-4 到 3e-4Transformer 块多的模型用低一点卷积为主的用高一点。调度用余弦退火加 warmupwarmup 设 5 到 10 个 epoch让注意力权重先稳定下来。from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR optimizer AdamW(model.parameters(), lr3e-4, weight_decay1e-4) warmup LinearLR(optimizer, start_factor0.01, total_iters5) cosine CosineAnnealingLR(optimizer, T_max95, eta_min1e-6) scheduler SequentialLR(optimizer, schedulers[warmup, cosine], milestones[5])warmup 的 start_factor 设 0.01 意味着从 3e-6 开始爬5 个 epoch 到 3e-4。如果跳过 warmupTransformer 块的自注意力在初期容易输出均匀分布梯度信号很弱收敛会慢很多。4.2 批次大小与显存权衡Synapse 2D 切片训练输入 224x224 或 256x256批次大小 8 到 16 比较常见。显存不够就减批次加梯度累积。我一般用 12GB 显存跑 256x256、批次 8如果 Transformer 块放在 32x32 特征图上显存大概占 10GB 左右。accumulation_steps 2 for i, (img, lbl) in enumerate(train_loader): pred model(img) loss criterion(pred, lbl) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()梯度累积等效于增大批次但 BatchNorm 的统计量还是按实际批次算的所以批次不要小于 4否则 BN 的均值方差估计不准。如果显存实在紧张可以把 BN 换成 GroupNorm对批次大小不敏感。4.3 训练轮数与早停策略Synapse 30 例数据2D 切片大概几千张训练 100 到 150 个 epoch 比较合适。太多会过拟合验证集 Dice 在 80 到 100 epoch 之间通常会到峰值然后缓慢下降。早停的 patience 设 15 到 20监控验证集的平均 Dice。best_dice 0 patience 20 counter 0 for epoch in range(150): train_one_epoch(model, train_loader, optimizer, criterion) val_dice evaluate(model, val_loader) scheduler.step() if val_dice best_dice: best_dice val_dice torch.save(model.state_dict(), best_model.pth) counter 0 else: counter 1 if counter patience: print(fEarly stop at epoch {epoch}) break注意验证集要按患者划分不能按切片随机划分。同一个患者的相邻切片非常相似如果随机划分验证集里会有大量和训练集几乎一样的切片Dice 会虚高 5 到 10 个点。这个坑我踩过后来改成按患者划分才看到真实性能。4.4 混合精度训练与显存优化混合精度用 torch.cuda.amp前向用 float16反向用 float32 更新权重。显存能省 30% 到 40%速度也能快 20% 左右。但要注意 Dice Loss 里的 softmax 在 float16 下可能溢出需要把 loss 计算放在 float32 里。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for img, lbl in train_loader: optimizer.zero_grad() with autocast(): pred model(img) loss criterion(pred.float(), lbl) # loss 用 float32 scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()GradScaler 会自动处理梯度缩放防止 float16 下梯度下溢。如果训练中出现 loss 变成 nan先检查是不是在 autocast 里算了 softmax 或 log这些操作对数值范围敏感。5. 避坑与排查Synapse 8 类分割里最容易翻车的 5 个地方5.1 小器官 Dice 为 0 或极低现象训练完一看胆囊、胰腺的 Dice 只有 0.1 到 0.3甚至有几例直接是 0。原因类别不平衡导致模型倾向于忽略小器官加上切片筛选时小器官的正样本被大量过滤模型见得太少。解决第一切片筛选的 min_organ_pixels 不要设太高确保每个器官都有足够正样本第二损失函数里给小器官加权重或者用 Focal Loss 替代 CE第三验证时单独看每个器官的 Dice不要只看平均值。我一般会把胆囊和胰腺的权重设到 3训练后 Dice 能到 0.5 到 0.6。5.2 验证集 Dice 虚高但测试集崩了现象验证集平均 Dice 0.82测试集只有 0.65。原因验证集和训练集按切片随机划分同一患者的相邻切片同时出现在两边模型实际上是在“背”患者。解决按患者划分数据集训练、验证、测试三者的患者 ID 完全不重叠。如果数据量太少至少保证验证集的患者不在训练集里。这个改动会让验证 Dice 掉几个点但测试集的表现才是真实的。5.3 训练 loss 震荡不收敛现象loss 在 0.5 到 1.5 之间来回跳几个 epoch 都不降。原因学习率太大或者 warmup 太短Transformer 块的注意力还没稳定就开始大更新。解决降低初始学习率到 1e-4延长 warmup 到 10 个 epoch检查 BatchNorm 的 momentum 是不是太大默认 0.1可以降到 0.01。另外确认输入归一化是不是正确HU 值没截断的话梯度会爆炸。5.4 显存溢出OOM现象训练到一半报 CUDA out of memory或者一开始就 OOM。原因Transformer 块放在太大的特征图上序列长度平方增长或者批次大小设太大。解决把 Transformer 块从 64x64 特征图移到 32x32 或 16x16用梯度累积替代大批次开启混合精度如果还不行用窗口注意力或把自注意力换成线性注意力。我一般会在编码器最后两层才加 Transformer前面纯卷积。5.5 推理时输入尺寸和训练不一致导致 Dice 下降现象训练时 256x256推理时用原始 512x512Dice 掉 10 个点以上。原因位置嵌入是按 256x256 学的换尺寸后位置编码对不上BatchNorm 的统计量也是按训练尺寸算的。解决推理时保持和训练一致的输入尺寸或者对位置嵌入做双线性插值。如果必须用原始尺寸把 BatchNorm 换成 GroupNorm位置嵌入改成正弦编码对尺寸不敏感。我一般固定 256x256推理时把原始切片 resize 到这个尺寸再 resize 回原尺寸算 Dice。6. 结果验证与进阶技巧per-class Dice 怎么看、模型怎么再提点6.1 用混淆矩阵定位问题器官平均 Dice 会掩盖很多问题。8 类分割里肝脏和脾脏的 Dice 通常能到 0.85 以上但胆囊和胰腺可能只有 0.5。我一般会画一个 9x9 的混淆矩阵看哪些类别之间互相混淆。常见的是脾脏和肝脏在边界处混胃和胰腺在胃壁附近混。from sklearn.metrics import confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def plot_confusion(preds, targets, num_classes9): preds_flat preds.flatten() targets_flat targets.flatten() cm confusion_matrix(targets_flat, preds_flat, labelslist(range(num_classes))) cm_norm cm.astype(float) / cm.sum(axis1, keepdimsTrue) sns.heatmap(cm_norm, annotTrue, fmt.2f, cmapBlues) plt.xlabel(Predicted) plt.ylabel(True) plt.show()看混淆矩阵的时候重点看对角线以外的值。如果脾脏的预测有 15% 跑到肝脏上说明这两个器官的特征区分度不够可以考虑在损失里加一个类间距离的约束或者用更深的编码器。6.2 后处理连通域过滤与形态学操作模型输出的小器官预测经常有零散的小连通域这些是假阳性。用连通域分析去掉小于一定体积的区域能提 1 到 2 个点的 Dice。from scipy.ndimage import label def remove_small_components(pred_mask, min_size50): cleaned np.zeros_like(pred_mask) for c in range(1, pred_mask.max() 1): mask (pred_mask c) labeled, num label(mask) for i in range(1, num 1): if np.sum(labeled i) min_size: cleaned[labeled i] c return cleanedmin_size 按器官来设胆囊可以设 30肝脏设 200。设太大小器官的真阳性会被误删设太小假阳性去不掉。我一般会在验证集上扫一遍 min_size看哪个值让 per-class Dice 最高。6.3 模型集成与测试时增强单模型 Dice 到瓶颈之后可以用集成再提一点。常见做法是训练 3 到 5 个不同初始化的模型推理时对 softmax 输出取平均。测试时增强TTA也可以用水平翻转 小角度旋转每个样本推理 4 到 8 次再平均。def tta_predict(model, img): preds [] # 原始 preds.append(torch.softmax(model(img), dim1)) # 水平翻转 img_flip torch.flip(img, dims[3]) pred_flip torch.softmax(model(img_flip), dim1) preds.append(torch.flip(pred_flip, dims[3])) # 平均 return torch.mean(torch.stack(preds), dim0)TTA 的代价是推理时间翻几倍如果做 8 次 TTA推理时间就是 8 倍。实际部署时要在精度和速度之间权衡。我一般只在最终提交或论文对比时用 TTA日常验证不用。6.4 我踩过的最大一个坑训练集和验证集按切片随机划分验证 Dice 0.84测试集 0.61。当时以为是模型泛化不行换了几个结构都没用。后来把患者 ID 列出来一看训练集和验证集里有 6 个患者是重叠的。改成按患者划分之后验证 Dice 掉到 0.72但测试集也是 0.70这才是真实水平。从那以后我拿到任何医学数据集第一件事就是确认划分单位是患者还是切片这个习惯帮我省了很多后悔药。希望帮到你。本文还有配套的精品资源点击获取