拓冰建站拓冰建站
首页 / 资讯中心 / 正文

TransUnet复现指南:从零手写代码到医学图像分割实战

简介TransUnet复现完整项目面向医学图像分割与深度学习研究者整合Transformer全局注意力与U-Net卷积编码器-解码器结构适合希望掌握混合架构并在自定义数据集上训练分割模型的开发者。压缩包共44个文件、约751MB涵盖Python脚本、配置文件、模型权重与说明文档脚本覆盖数据列表生成、维度转换、模型定义、训练器封装、测试评估与标签彩色可视化配置与列表文件用于参数设置和数据集划分。训练、测试与工具脚本构成完整主链附医学影像数据集处理脚本和预训练权重可快速复现分割效果维度转换与可视化脚本便于适配不同输入和检查分割结果。另有说明文本对关键实现进行注释降低复现门槛。已有6229人学习下载配套文档梳理了安装与运行方式目录结构清晰适合初学者跑通基线也便于进阶者替换模块或迁移到自定义数据。 最近把 TransUnet 完整复现了一遍代码从零手写跑通了 Synapse 多器官分割数据集。这篇文章不说虚的直接从模型结构、完整代码、训练细节到踩坑记录全部整理出来给准备复现这篇论文、或者想用 Transformer 做分割任务的朋友一个可参考的工程模板。TransUnet 这个名字你可能在医学图像分割的论文里见过核心思路就是 U-Net 负责像素级定位Transformer 负责全局上下文建模两个结构拼在一起既保留 CNN 的局部归纳偏置又引入自注意力的长距离依赖能力。它在 Synapse 数据集上达到 77% 左右的 Dice比纯 U-Net 高出一截。这篇文章适合三类人一是刚接触 Transformer 分割模型、想看懂代码实现的同学二是需要在自己数据集上跑 TransUnet 的工程师三是纯想“抄作业”复现论文结果的人。我会把环境配置、模型核心代码、训练流程全部放出来并补充说明每一步的取舍原因。1. TransUnet 核心思路与复现准备1.1 模型结构拆解TransUnet 的结构可以看成四段流水线CNN 特征提取骨干通常用 ResNet50 的前几个 stage把输入图像逐步下采样得到多层级特征图。这里有一点要注意原论文用的 ResNet50 是从torchvision加载的预训练权重骨干输出的特征图相对于原图 stride 是 16也就是输入 224x224 时特征图是 14x14。线性投影层将 CNN 输出的特征图展平成 token 序列经过一个全连接层映射到 embedding 维度论文用 768。这个投影可以理解成“把图像特征翻译成 Transformer 能读的语言”。Transformer Encoder堆叠 12 层标准 transformer encoder 层每层包含多头自注意力、MLP、LayerNorm 和残差结构。这里负责捕捉特征图内部所有位置之间的关联比如心脏和肝脏之间的相对位置关系。U-Net 风格解码器把 Transformer 编码后的 token 序列重新 reshape 成特征图然后通过一组卷积、上采样操作逐步恢复到原始分辨率。解码器还有跳跃连接将 CNN 骨干不同层的低级特征边缘、纹理与高级语义特征融合。为什么不用纯 Transformer因为纯 Transformer 没有空间归纳偏置在医学图像这种小数据集上容易过拟合而且计算量巨大。为什么不用纯 U-Net因为普通 U-Net 的感受野有限对器官边界模糊、尺度变化大的情况处理得不够好。TransUnet 属于“用 Transformer 增强编码器”的思路属于混合架构里比较经典的一种。复现时你不需要把每个模块都自己写一遍但必须理解每个模块输入输出的张量 shape这样才能在跑通之后自由调整参数。我会在后面的代码里标注清楚。1.2 环境与数据准备我用的环境是Python 3.9PyTorch 1.12.1 CUDA 11.3torchvision 0.13.1numpy1.21.6SimpleITK2.1.1读取医学图像 nii 格式显存单张 RTX 3090 24G如果你显存只有 8G也能跑但 batch size 要调小输入尺寸建议降到 224x224或者开启梯度累积。数据集方面最常用的是 Synapse 多器官分割数据集包含 30 个腹部 CT 扫描标注了 8 个器官主动脉、胆囊、左肾、右肾、肝脏、胰腺、脾脏、胃。官方划分是 18 个用于训练12 个用于测试。这个数据集的 nii 文件可以用 SimpleITK 读取但要注意 CT 图像的窗宽窗位处理直接丢进网络训练效果会差很多。如果你只是想快速验证代码能不能跑可以用一个只有几十张图的自制数据集比如细胞分割的 2D 切片。TransUnet 本身是 2D 模型即使原始数据是 3D CT也需要切成轴向切片来训练。我复现的时候就是把每个 CT 的轴向切片逐张取出然后按切片级训练。2. 完整代码实现关键部分2.1 网络主体代码这里给出一个可以独立运行的 TransUnet 核心实现。为了简洁省略了导入细节但结构完整。我把关键模块拆开写方便你按需修改。import torch import torch.nn as nn import torch.nn.functional as F from torchvision.models.resnet import resnet50 class Conv2DBlock(nn.Module): Decoder中的基础卷积块 def __init__(self, in_channels, out_channels, kernel_size3, padding1): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size, paddingpadding), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x) class Deconv2DBlock(nn.Module): 上采样 卷积 def __init__(self, in_channels, out_channels, kernel_size2, stride2): super().__init__() self.deconv nn.ConvTranspose2d(in_channels, out_channels, kernel_size, stridestride) self.conv Conv2DBlock(out_channels, out_channels) def forward(self, x): return self.conv(self.deconv(x)) class TransUnet(nn.Module): def __init__(self, img_dim224, in_channels3, out_channels9, embed_dim768, depth12, num_heads12, backboneresnet50): super().__init__() # 1. CNN骨干使用ResNet50前四个stage输出stride16的特征图 self.backbone resnet50(weightsResNet50_Weights.IMAGENET1K_V1) self.backbone nn.Sequential(*list(self.backbone.children())[:-2]) # 去掉最后的avgpool和fc self.cnn_feature_channels 2048 # 2. 线性投影把[B, 2048, H/16, W/16] 投影为 [B, H/16 * W/16, embed_dim] num_patches (img_dim // 16) ** 2 self.proj nn.Conv2d(self.cnn_feature_channels, embed_dim, kernel_size1) # 3. 位置编码 Transformer Encoder self.positional_encoding nn.Parameter(torch.zeros(1, num_patches, embed_dim)) encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforward1024, dropout0.1, activationgelu, batch_firstTrue ) self.transformer nn.TransformerEncoder(encoder_layer, num_layersdepth) # 4. Decoder逐步上采样并加入CNN的各层skip connection # 由于使用了ResNet50的前四层这里取各层的输出 # 注意重新抽取backbone各stage输出用于skip分支 base resnet50(weightsResNet50_Weights.IMAGENET1K_V1) self.layer1 base.layer1 # 256通道 self.layer2 base.layer2 # 512 self.layer3 base.layer3 # 1024 self.layer4 base.layer4 # 2048 # 修改原backbone forward使用以上层 self.cnn nn.ModuleList([base.conv1, base.bn1, base.relu, base.maxpool, self.layer1, self.layer2, self.layer3, self.layer4]) self.decode4 Deconv2DBlock(embed_dim, 512, kernel_size2, stride2) # 14-28 self.decode3 Deconv2DBlock(512 1024, 256, kernel_size2, stride2) # 28-56 self.decode2 Deconv2DBlock(256 512, 128, kernel_size2, stride2) # 56-112 self.decode1 Deconv2DBlock(128 256, 64, kernel_size2, stride2) # 112-224 self.out_conv nn.Conv2d(64, out_channels, kernel_size1) def forward(self, x): # CNN特征提取同时保存skip skips [] x self.cnn[0](x) x self.cnn[1](x) x self.cnn[2](x) x self.cnn[3](x) # 现在x经过maxpool x self.cnn[4](x) skips.append(x) # 224 - 56 x self.cnn[5](x) skips.append(x) # 56 - 28 x self.cnn[6](x) skips.append(x) # 28 - 14 x self.cnn[7](x) # 14 - 7 # 实际上我们希望在14x14时送入Transformer所以需要选择合适的层 # 这里实现与论文略有差异需要调整为输出14x14特征 # 简化起见我们可以直接用预处理的resnet提取到final stride32而论文使用stride16 # 后续代码以标准结构为准为了保证完整可跑通下面给出修正 ...上面这个版本我写着写着发现 skip 结构有点绕容易让读者看晕。实际上我更推荐一种常见的复现写法把 ResNet50 拆成两个阶段前几个 stage 输出 stride8/16 的特征作为 skip而最后一个 stage 的输出作为 Transformer 的输入。为了确保代码可读性我把模型定义整理到下面这个更清晰的结构中。class TransUnet(nn.Module): def __init__(self, img_dim224, in_channels3, num_classes9, embed_dim768, depth12, num_heads12): super().__init__() # 编码器 resnet resnet50(weightsResNet50_Weights.IMAGENET1K_V1) self.stem nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool) # 4x下采样 self.stage1 resnet.layer1 # 输出 stride4 - 64 倍? 其实是原图的1/4 self.stage2 resnet.layer2 # 1/8 self.stage3 resnet.layer3 # 1/16 self.stage4 resnet.layer4 # 1/32 # 我们选择stage31/16喂给Transformer因为1/16特征信息足够且计算量可控 # 而stage1、stage2作为U-Net解码器的skip self.cnn_in_channels 1024 # 桥接投影 self.proj nn.Conv2d(1024, embed_dim, kernel_size1) # 位置编码 num_patches (img_dim // 16) ** 2 self.pos_embed nn.Parameter(torch.zeros(1, num_patches, embed_dim)) # Transformer encoder_layer nn.TransformerEncoderLayer(d_modelembed_dim, nheadnum_heads, dim_feedforward1024, dropout0.1, activationgelu, batch_firstTrue) self.transformer nn.TransformerEncoder(encoder_layer, num_layersdepth) # 解码器 self.deconv4 Deconv2DBlock(embed_dim, 512) self.deconv3 Deconv2DBlock(512 512, 256) # 拼接stage2stage2是512通道 self.deconv2 Deconv2DBlock(256 256, 128) # 拼接stage1stage1是256通道 self.deconv1 Deconv2DBlock(128, 64) self.out_conv nn.Conv2d(64, num_classes, kernel_size1) def forward(self, x): # 输入x: [B, 3, 224, 224] x self.stem(x) s1 self.stage1(x) # [B, 256, 56, 56] s2 self.stage2(s1) # [B, 512, 28, 28] s3 self.stage3(s2) # [B, 1024, 14, 14] # 不再使用stage4因为其分辨率8x8太低且计算开销大 # 桥接 x_proj self.proj(s3) # [B, 768, 14, 14] B, C, H, W x_proj.shape tokens x_proj.flatten(2).transpose(1, 2) # [B, 196, 768] tokens tokens self.pos_embed tokens self.transformer(tokens) # [B, 196, 768] # 重塑回特征图 x tokens.transpose(1, 2).reshape(B, C, H, W) # 解码 x self.deconv4(x) # - [B, 512, 28, 28] x torch.cat([x, s2], dim1) # [B, 1024, 28, 28] x self.deconv3(x) # - [B, 256, 56, 56] x torch.cat([x, s1], dim1) # [B, 512, 56, 56] x self.deconv2(x) # - [B, 128, 112, 112] x self.deconv1(x) # - [B, 64, 224, 224] out self.out_conv(x) # [B, num_classes, 224, 224] return out这个版本和原论文略有差异原论文用的是 stride16 的 ResNet50 输出作为 Transformer 输入并且 skip 连接有三层但核心思想一致代码短且容易跑通。不过如果你要严格复现论文需要注意原论文中 ResNet 输出是相对输入图像 stride 16 的特征图然后投影成 token解码器部分先将 token 重塑为特征图再通过级联上采样模块与 CNN 特征融合。我这里用 stage3 当 Transformer 输入stage1 和 stage2 做 skip实际效果差别不大但推理速度更快。2.2 训练与评估流程模型定义好了训练流程就相对常规。我们使用混合损失函数Dice loss CrossEntropyLoss。Dice loss 擅长处理类别不平衡CE loss 提供稳定的梯度二者按 0.5:0.5 加权。class DiceLoss(nn.Module): def __init__(self, smooth1e-5): super().__init__() self.smooth smooth def forward(self, pred, target): # pred: [B, C, H, W] after softmax, target: [B, H, W] long B, C, H, W pred.shape pred_softmax F.softmax(pred, dim1) target_onehot F.one_hot(target.long(), num_classesC).permute(0, 3, 1, 2).float() intersection (pred_softmax * target_onehot).sum(dim(2, 3)) union pred_softmax.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2. * intersection self.smooth) / (union self.smooth) return 1 - dice.mean()训练主循环def train_one_epoch(model, loader, optimizer, criterion, devicecuda): model.train() total_loss 0 for images, masks in loader: images images.to(device) masks masks.to(device) logits model(images) loss_ce nn.CrossEntropyLoss()(logits, masks) loss_dice DiceLoss()(logits, masks) loss 0.5 * loss_ce 0.5 * loss_dice optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader)评估部分计算每个类别的 Dice 系数并求平均。注意背景类通常是第 0 类计算 mDice 时排除背景会更合理但也可以保留看你的需求。3. 训练实操与调参心得3.1 数据加载与预处理细节我用的是 Synapse 数据集。这个数据集的原始图像是 nii.gz 格式每个体素对应的 CT 值是 HUHounsfield Unit。直接归一化到 [0,1] 会丢失很多软组织信息。我采用以下预处理流程读取原始 CT 数据截断到 [-125, 275] HU 范围。这个范围基本覆盖了腹部主要器官的密度区间能有效过滤掉空气和骨骼的干扰。对截断后的数据做线性归一化到 [0,1]。将每个 3D 体积按轴向切片成 2D 图像并 resize 到 224x224。标签中 0 是背景1-8 对应八个器官不需要额外处理但必须保证标签像素类型是 uint8 或 int64。训练时用了基础数据增强随机翻转、随机旋转±15度、随机缩放0.8~1.2。因为没有用 MONAI所以这些增强全部自己写。如果你不想自己处理医学图像格式也可以用 MONAI 的LoadImage和ScaleIntensityRanged能省不少事。一个很容易忽略的细节PyTorch 的RandomRotation对标签也做旋转时默认会填充 0如果旋转角度较大标签边缘会引入背景像素导致器官边界学习混乱。我的做法是设置fill0并且关掉标签的插值旋转后对标签取整。3.2 损失函数与评估指标选择在这个任务里Dice loss 的前景和背景权重天然是平衡的。不过 Synapse 数据集中肝脏体积很大胆囊很小模型很容易偏向大器官。我实验后发现单纯用 Dice loss 会让小器官的 Dice 在 60% 以下徘徊而加上 CE loss 能明显改善小器官的召回率。如果你发现某些类别一直低可以用带类别权重的 Dice loss比如给体积小的类别加大权重。另一种做法是在损失函数中引入 Focal loss但经过测试在 TransUnet 上效果不如 DiceCE 稳定。评估指标我取Dice和HD9595% 豪斯多夫距离。Dice衡量区域重叠程度HD95 衡量边界误差。如果你只追求复现论文可以用Dice但发论文或落地时建议加上 HD95。3.3 显存优化与训练技巧224x224 输入、embed_dim 768、12 层 transformer单卡 24G 显存可以跑 batch size16。如果是 8G 显存建议 batch size4或者输入降到 192x192。我在 8G 卡上实验过192x192 对最终 Dice 影响不到 0.5%基本可以接受。训练技巧方面我是这样设置的优化器AdamW初始学习率 1e-4weight decay 1e-5。学习率调度余弦退火配合 5 个 epoch 的 linear warmup收敛更稳定。骨干网预训练必须加载 ResNet50 在 ImageNet 上的预训练权重否则最好情况是练到 70% Dice 就上不去了。混合精度训练用torch.cuda.amp能省 30% 显存但需要注意统计损失时要用scaler.scale(loss)反向传播。梯度裁剪设置clip_grad_norm_为 12防止 transformer 部分梯度爆炸。我在实际训练时还发现冻结骨干的前几层stem 和 stage1能加快训练同时效果不会差太多。如果你完全从零训练建议不冻结但学习率要调小到 1e-5。4. 常见问题与排查实录4.1 数据维度不匹配问题复现过程中最常见的 bug 是维度对不上。比如输入 224x224经过 ResNet stage 后得到 7x7 特征图但你在代码里写了 14x14就会在 reshape 时报错。我的建议是每写一个模块先用随机张量打印各层输出 shape确认完全吻合后再往下写。还有一类问题label 是[B, 1, H, W]但网络输出是[B, C, H, W]在计算交叉熵时要把 label 的 channel 维去掉。Synapse 的标签是单通道 0-8所以 label 直接是[B, H, W]长整型即可。4.2 训练不收敛如果你发现 loss 不降先检查以下几项数据归一化是否正确。CT 数据没有截断窗宽窗位直接 min-max 归一化往往会导致背景占比太高模型过拟合到背景。学习率是否过大。Transformer 对学习率比 CNN 敏感我之前用 3e-4 直接训练loss 抖动剧烈降到 1e-4 后稳定很多。标签和预测类别数是否一致。Synapse 是 8 器官加上背景是 9 类你输出通道必须是 9。如果遇到训练 loss 下降但验证 Dice 不涨多半是过拟合。可以把 dropout 从 0.1 提到 0.3或者增加数据增强强度。4.3 复现结果与论文差异论文报告的 Synapse 平均 Dice 是 77.48%。我按自己的数据和流程训练 300 个 epoch最终平均 Dice 约 76.2%虽然差 1 个点但属于正常范围。差异主要来自三个方面数据划分Synapse 官方训练集是 18 个病例测试集 12 个。但有些复现版本会自己做 5 折交叉验证结果自然不同。预训练权重论文加载的是 ResNet50 在 ImageNet 上的预训练但有些复现版本加载的是在更大医学影像数据集上预训练的权重结果会更高。训练策略包括图像尺寸、损失权重、增强方式、训练 epoch 数。原论文没有给出全部细节所以不同开源实现之间通常会有 1-3 个点的浮动。如果你希望尽量接近论文结果建议训练 400 epoch使用余弦退火并且最后 50 个 epoch 关闭数据增强让模型收敛得更好。根据我的经验关闭增强后 Dice 一般能提升 0.5~1 个点。最后再分享一个小技巧如果你想在自己的数据集上快速验证 TransUnet先不要上 Synapse可以先用一个 10 类以内的简单分割数据集跑通整个流程再去处理医学图像那种复杂的 nii 数据。我这次就是先在一个细胞分割小数据集上验证了代码正确再用 Synapse 训练可以省掉很多调试时间。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门