
简介一套完整的TransUnet复现工程面向从事医学图像分割研究或入门TransformerU-Net架构的开发者帮助理解并落地这一经典模型。压缩包共44个文件约751MB包含14个Python脚本、6个pyc编译文件、2个PyTorch权重文件以及7个xml配置、9个txt说明和2个Markdown文档等。Python脚本覆盖了模型定义如vit_seg_modeling.py、训练主程序train.py/trainer.py、测试评估test.py、数据预处理与列表生成make_list_file.py/To_2d.py/To_3d.py以及分割结果彩色可视化show_label_to_color.py配合README和说明文档构成从数据到推理的完整链路。已有6229人学习下载。通过这套资源使用者可以对照代码理解Transformer全局注意力如何与U-Net跳跃连接结合利用预训练权重快速验证分割效果并参考实现细节调整以适配自定义的医学图像数据集。 做医学图像分割的同学应该都绕不开TransUNet。这个模型在Synapse、ACDC这些公开数据集上被反复当作baseline2021年发表在Medical Image Analysis上到现在依然是“CNN提特征 Transformer建全局关系 U型解码”这种混合架构最直接的范本。我这次把复现过程中整理的一套完整代码和实现说明放出来代码不依赖vit_pytorch这类第三方封装用PyTorch原生组件就能跑适合正在搭baseline、或者想改到自己数据集上做实验的人参考。这套东西我前前后后踩了不少坑尤其是位置编码的尺寸对齐、官方预训练权重的key映射、mask做resize时把标签插值“插坏”这几个问题几乎每个人都会遇到。下面先讲清楚结构上的关键点再给完整代码最后把排查思路整理成链路方便你在自己的数据上报错时按图索骥。1. 复现前必须先想清楚的三个问题1.1 混合Encoder到底改了U-Net的什么U-Net本身是纯卷积结构Encoder部分通过逐级池化拿到多尺度的局部特征再通过skip connection把浅层细节传给Decoder。CNN的优势是局部归纳偏置强小样本也能学得不错但代价是感受野受限对器官边界模糊、对比度低的区域全局上下文建模能力不够。TransUNet的思路不是把U-Net推翻而是把U-Net最底层的瓶颈特征图交给Transformer处理。具体来说先用CNN把输入图像下采样到1/16分辨率得到 feature map然后切成patch序列加上位置编码进入多层Transformer Encoder。Transformer的自注意力机制能建模任意两个位置之间的长距离依赖弥补CNN只看局部窗口的不足。处理完的序列再还原成二维特征图交给Decoder逐级上采样并在每一级和CNN中间层特征做拼接保留精细的边界信息。所以它本质上是一个“混合Encoder U型Decoder”CNN负责低层细节Transformer负责全局语义Decoder负责把语义映射回像素空间。这个设计思路在后来的UNETR、Swin-Unet里都能看到影子复现它相当于把这一整条技术路线吃透。1.2 官方仓库代码为什么不能直接拿来用官方仓库代码能用但直接搬到自己项目里会很痛苦。原因有三点第一官方代码的配置文件、数据集预处理、模型封装耦合得很深我当年第一次跑的时候光整理Synapse数据集的nii.gz文件和归一化逻辑就花了大半天。第二主干网络支持R50-ViT和ViT两种模式权重分支名有好几套换预训练权重时经常出现state_dict key对不上。第三官方实现里混着大量实验遗留参数比如auxiliary loss分支、skip connection数量可调等等对于只想快速搭一个稳定baseline的人来说这些反而是负担。因此我这次的复现目标很明确做一版结构清晰、无第三方依赖、拿过来就能改的轻量级实现。模型保留TransUNet的核心思想——CNN下采样、Transformer编码、U型解码和skip connection但代码精简到四个文件训练数据和预测路径都直接可替换。1.3 本次复现的运行环境与预设先说明我这边跑通的环境方便你对照检查Python 3.8 / 3.10都可以PyTorch 1.12以上建议2.0torchvision只在需要加载官方预训练权重时用到基础版本不需要显卡显存建议8GB以上我用的是RTX 3060 12GBbatch size开到8没问题数据集用标准png格式的图像和mask图像建议统一到224x224下面所有代码都直接基于这套环境编写。如果你的PyTorch版本比较老注意把nn.MultiheadAttention的batch_firstTrue参数确认一下老版本不认识这个参数需要手动把输入维度换成seq_len, batch, embed_dim。2. 完整可运行的轻量版复现代码2.1 项目文件结构代码按功能拆成四个文件不搞花活transunet_repro/ ├── model.py # 网络结构定义 ├── dataset.py # 数据加载器 ├── train.py # 训练脚本 ├── predict.py # 推理脚本 ├── data/ │ ├── images/ # 训练原图png格式 │ └── masks/ # 对应标签png格式二值或灰度 └── checkpoints/ # 模型权重保存目录data/images和data/masks下的文件名一一对应这是最简单的组织方式。如果你的数据是nii.gz格式先转成png或npy再做训练这里面少踩很多格式转换的坑。2.2 model.py核心网络实现模型部分我按流程拆成五个模块CNN下采样Encoder、Patch嵌入、Transformer Block、Decoder、整体组装。建议你按顺序读代码不要直接跳到最后看完整类。# model.py import torch import torch.nn as nn class CNNEncoder(nn.Module): CNN下采样Encoder输出四层特征分别对应1/2、1/4、1/8、1/16分辨率 def __init__(self, in_ch3): super().__init__() self.stage1 nn.Sequential( nn.Conv2d(in_ch, 64, 3, stride2, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.Conv2d(64, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), ) # 1/2 self.stage2 nn.Sequential( nn.Conv2d(64, 128, 3, stride2, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.Conv2d(128, 128, 3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), ) # 1/4 self.stage3 nn.Sequential( nn.Conv2d(128, 256, 3, stride2, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.Conv2d(256, 256, 3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), ) # 1/8 self.stage4 nn.Sequential( nn.Conv2d(256, 512, 3, stride2, padding1), nn.BatchNorm2d(512), nn.ReLU(inplaceTrue), nn.Conv2d(512, 512, 3, padding1), nn.BatchNorm2d(512), nn.ReLU(inplaceTrue), ) # 1/16 def forward(self, x): s1 self.stage1(x) s2 self.stage2(s1) s3 self.stage3(s2) s4 self.stage4(s3) return s1, s2, s3, s4 class PatchEmbed(nn.Module): 将CNN输出的1/16特征图切成token序列 def __init__(self, in_ch512, embed_dim384, patch_size1): super().__init__() self.proj nn.Conv2d(in_ch, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # B, embed_dim, H, W B, C, H, W x.shape x x.flatten(2).transpose(1, 2) # B, H*W, embed_dim return x, H, W这里patch_size1是刻意为之。CNN已经完成了16倍下采样1/16分辨率特征图上的每个像素就等同于原始输入的16x16 patch这个写法既保留了TransUNet的patch思想又省掉了一次reshape的麻烦。接下来是Transformer部分class TransformerBlock(nn.Module): 标准Transformer Encoder Block def __init__(self, dim, num_heads, mlp_ratio4., 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_norm self.norm1(x) x x self.attn(x_norm, x_norm, x_norm)[0] x x self.mlp(self.norm2(x)) return x class DecoderBlock(nn.Module): U型Decoder上采样后与skip特征拼接 def __init__(self, in_ch, skip_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, out_ch, kernel_size2, stride2) self.conv nn.Sequential( nn.Conv2d(out_ch skip_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), ) def forward(self, x, skip): x self.up(x) x torch.cat([x, skip], dim1) return self.conv(x)TransformerBlock里有一个小细节先把self.norm1(x)存下来再传进attention而不是在forward里写三次self.norm1(x)。虽然结果等价但少做两次LayerNorm前向和反向都会快一点。DecoderBlock的上采样统一用ConvTranspose2dkernel_size2、stride2刚好把分辨率翻倍和skip在通道维拼接再接两层卷积融合。最后组装成完整的TransUNetclass TransUNet(nn.Module): def __init__(self, in_ch3, num_classes1, patch_size1, embed_dim384, depth6, num_heads6, dropout0.1): super().__init__() # CNN编码器 self.encoder CNNEncoder(in_ch) # 特征图 - token self.patch_embed PatchEmbed(in_ch512, embed_dimembed_dim, patch_sizepatch_size) self.pos_embed nn.Parameter(torch.zeros(1, (224 // 16) ** 2, embed_dim)) nn.init.trunc_normal_(self.pos_embed, std0.02) # Transformer编码器 self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads, dropoutdropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) # U型解码器 self.decoder1 DecoderBlock(embed_dim, 512, 256) # 1/16 - 1/8 self.decoder2 DecoderBlock(256, 256, 128) # 1/8 - 1/4 self.decoder3 DecoderBlock(128, 128, 64) # 1/4 - 1/2 self.decoder4 DecoderBlock(64, 64, 32) # 1/2 - 1 self.head nn.Conv2d(32, num_classes, kernel_size1) def forward(self, x): B, _, H, W x.shape assert H % 16 0 and W % 16 0, 输入尺寸必须是16的倍数 # CNN编码 s1, s2, s3, s4 self.encoder(x) # token化 Transformer x, fh, fw self.patch_embed(s4) x x self.pos_embed[:, :x.size(1), :] for blk in self.blocks: x blk(x) x self.norm(x) # 还原为二维特征图 x x.transpose(1, 2).view(B, -1, fh, fw) # 解码 x self.decoder1(x, s3) x self.decoder2(x, s2) x self.decoder3(x, s1) x self.decoder4(x, x) # 这里注意最后一级没有额外的skip x self.head(x) return x最后一级decoder这里我留了一个容易混淆的地方self.decoder4(x, x)。因为decoder3的输出和decoder4的输入本身是同分辨率不需要skip拼接但DecoderBlock的forward必须接收两个参数所以直接传自己。严格来说这不太优雅但能少定义一个“没有skip的DecoderBlock”代码量更小。你自己改的时候可以单独写一个DecoderBlockNoSkip类逻辑更清晰。关于位置编码我初始化的是224分辨率、patch_size16对应的196个token。如果你训练时改成512x512token数变成1024直接相加就会报维度不匹配。后面第4节会专门讲这个问题。2.3 dataset.py和train.py数据加载与训练数据加载器的关键点是mask的resize必须用最近邻插值。原因很直接mask是离散标签线性插值会把0和1插成0.3、0.7这种灰度等于给模型喂了错误标签。# dataset.py import os from glob import glob import cv2 import numpy as np import torch from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, img_dir, mask_dir, size(224, 224)): self.img_paths sorted(glob(os.path.join(img_dir, *.png))) self.mask_paths sorted(glob(os.path.join(mask_dir, *.png))) assert len(self.img_paths) len(self.mask_paths), 图片和标签数量不一致 self.size size def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img cv2.imread(self.img_paths[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, self.size) img img.astype(np.float32) / 255.0 img torch.from_numpy(img).permute(2, 0, 1) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.resize(mask, self.size, interpolationcv2.INTER_NEAREST) mask (mask 0).astype(np.float32) # 二分类问题标签转为0/1 mask torch.from_numpy(mask).unsqueeze(0) return img, mask训练脚本同样保持精简把BCE和Dice损失组合在一起# train.py import os import torch import torch.nn as nn from torch.utils.data import DataLoader from model import TransUNet from dataset import SegmentationDataset def dice_loss(pred, target): pred torch.sigmoid(pred) smooth 1.0 intersection (pred * target).sum() return 1 - (2.0 * intersection smooth) / (pred.sum() target.sum() smooth) def main(): os.makedirs(checkpoints, exist_okTrue) device cuda if torch.cuda.is_available() else cpu dataset SegmentationDataset(data/images, data/masks) dataloader DataLoader(dataset, batch_size8, shuffleTrue, num_workers4) model TransUNet(in_ch3, num_classes1).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) criterion nn.BCEWithLogitsLoss() for epoch in range(100): model.train() total_loss 0.0 for imgs, masks in dataloader: imgs, masks imgs.to(device), masks.to(device) logits model(imgs) loss criterion(logits, masks) dice_loss(logits, masks) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() print(fEpoch {epoch 1}, Loss: {total_loss / len(dataloader):.4f}) torch.save(model.state_dict(), checkpoints/last.pth) if __name__ __main__: main()BCE Dice这个组合是医学图像分割里最常用的训练目标。BCE从像素角度独立判断每个点对不对Dice从区域角度衡量预测和真实标签的重叠程度。两个目标叠加既能保证像素级精度又能缓解前景和背景像素数量严重不平衡的问题。2.4 predict.py推理与保存结果推理脚本最容易忽略的是model.eval()和torch.no_grad()。不加这两个测试时的Dropout和BatchNorm行为和训练不一致推理结果会抖动。# predict.py import torch import cv2 import numpy as np from model import TransUNet def main(): device cuda if torch.cuda.is_available() else cpu model TransUNet(in_ch3, num_classes1).to(device) model.load_state_dict(torch.load(checkpoints/last.pth, map_locationdevice)) model.eval() img cv2.imread(data/images/0000.png) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, (224, 224)) img img.astype(np.float32) / 255.0 x torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0).to(device) with torch.no_grad(): logit model(x) prob torch.sigmoid(logit).cpu().numpy()[0, 0] mask (prob 0.5).astype(np.uint8) * 255 cv2.imwrite(output.png, mask) if __name__ __main__: main()如果你要做多类别分割把模型输出的num_classes改成类别数head从1通道变成N通道训练时mask做one-hot损失函数换成CrossEntropyLoss。二分类先跑通再往多类扩展这个路径最稳。3. 这几个关键参数为什么这么定3.1 输入尺寸与patch_size的绑定关系Transformer的self-attention不限制序列长度限制长度的是位置编码。我的实现里pos_embed维度是(1, 196, 384)对应224分辨率除以16。只要输入不是224196和实际token数就对不上。如果你执意换输入尺寸正确做法是对位置编码做双线性插值import torch.nn.functional as F def resize_pos_embed(pos_embed, new_h, new_w): N, C pos_embed.shape[1], pos_embed.shape[2] old_h old_w int(N ** 0.5) pos_embed pos_embed.reshape(1, old_h, old_w, C).permute(0, 3, 1, 2) pos_embed F.interpolate(pos_embed, size(new_h, new_w), modebilinear, align_cornersFalse) pos_embed pos_embed.permute(0, 2, 3, 1).flatten(1, 2) return pos_embed但我的建议是除非有明确理由否则固定输入尺寸。模型是你设计的没有必要让输入尺寸成为变量。3.2 损失函数为什么选BCEDice医学图像分割里“背景远多于前景”是常态。以肝脏分割为例肝脏区域通常只占整个切片的10%到20%如果只用BCE模型很容易把全部像素预测为背景因为这样loss已经很低了。Dice系数天然对类别不平衡不敏感它计算的是集合重叠度直接把“预测区域和真实区域的相似程度”作为优化目标。所以把两者加起来让模型既要像素准确又要区域准确训练过程会更稳。3.3 优化器、学习率与批大小的选择代码里用的AdamW初始学习率1e-4。这是ViT类模型训练里的常见配置我一开始用的是1e-3结果loss在0.6附近震荡了很久降到1e-4之后才开始稳定下降。Transformer对学习率比纯CNN敏感如果你换成SGD学习率还需要重新调。batch size我设8这是在12GB显存下比较舒服的值。如果你的显卡显存不够优先降batch size而不是降输入分辨率。因为输入分辨率决定位置编码和token数量改了分辨率要连带改pos_embed牵一发动全身。另外注意源码里DataLoader(num_workers4)。Windows环境下如果多进程数据加载报错可以先改成num_workers0排查不是代码逻辑的问题。4. 复现中最容易踩的坑与排查链路4.1 位置编码尺寸错位最典型的RuntimeError现象模型forward时抛错报错内容类似The size of tensor a (196) must match the size of tensor b (200) at non-singleton dimension 1或者index out of range in self。原因输入分辨率不是patch_size的整数倍导致实际token数和pos_embed的数量对不上。我的实现里pos_embed是硬编码224对应的196个token你输入512x512时token数变成1024自然炸掉。排查步骤输入模型前打印x.shape确认H和W是不是16的倍数。打印model.pos_embed.shape[1]和H // 16 * W // 16对比。如果不等要么把输入统一resize到224要么用上面的resize_pos_embed插值后重新赋值。这一步的位置编码问题排查思路同时适用于你在别的Transformer模型里改动输入尺寸的场景。4.2 官方预训练权重的key对不上现象加载官方权重时提示Missing key(s) in state_dict或Unexpected key(s)权重加载失败但代码没报错训练出来的效果却很差。原因官方仓库同时支持R50-ViT和ViT两种主干权重的key命名不一样。比如R50-ViT的patch_embedding层key是vit.embeddings.patch_embeddings.projection.weight而纯ViT的key是conv1.weight。你自己改网络结构后key自然对不上。排查步骤别直接load先torch.load(xxx.pth, map_locationcpu)把state_dict打印出来看。和你的模型model.state_dict().keys()逐项对比找差异。如果只差一两个模块名写一个简单的映射函数替换key如果差异很大说明网络结构本身不一致检查你的实现里patch embedding和encoder的层数是否和官方一致。这个坑特别隐蔽它不会让你的程序崩溃只会让你的模型精度悄悄变差。训练之前一定要手动验证预训练权重能正确加载打印一行“load success”再继续。4.3 mask resize导致标签“虚化”现象训练loss下降很快Dice却不涨跑出来的预测图边缘发虚或者mask里有大量灰色过渡区域。原因cv2.resize(mask, size)默认使用线性插值mask里的0和1经过插值变成了0.3、0.7这种小数标签在训练时被当成了回归目标。虽然BCE也能吃小数标签但模型学到的东西已经不是“分割”而是“模糊的灰度估计”了。解决resize mask时必须加interpolationcv2.INTER_NEAREST。无论cv2还是PIL处理标签图时都要用最近邻插值这是个通用原则不只在TransUNet里成立。另外还要注意如果原始mask是灰度值0和255需要先mask 0转成0/1如果是多类别灰度图每种类别对应一个固定灰度值那你需要按类别生成one-hot标签不能直接mask 0。4.4 loss不降和显存不够两条高频问题链路loss不降先别急着改学习率。有一个固定的排查顺序打印logits的min、max、mean看有没有NaN。有NaN通常是梯度爆炸加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)同时把学习率降一个数量级。抽16张图在单batch上调到过拟合。如果小批量loss能降到接近0说明模型和数据加载没问题剩下的问题是训练策略比如学习率太大或数据增强太强。如果小批量也降不动检查标签是否有问题比如mask全是0或全没做归一化。显存不够最常见的是CUDA out of memory。除降batch size外可以尝试混合精度训练PyTorch 2.0之后用torch.autocast(cuda, dtypetorch.float16)包住前向和loss计算显存能省接近一半训练速度也更快。也可以用torch.utils.checkpoint对Transformer层做梯度检查点以时间换显存。但最简单的还是把embed_dim从384降到256depth从6降到4效果差距不大显存压力小很多。4.5 数据加载中的线程和路径坑补充一个很常见但很容易忽略的问题Windows下DataLoader(num_workers0)偶尔会卡死或报BrokenPipeError。这不是代码逻辑问题而是Windows的spawn方式和Linux的fork不同建议Windows用户先num_workers0等程序完全跑通再尝试调高。数据路径方面不要在代码里写绝对路径把data/images这种相对路径和项目根目录绑定换机器迁移时降低很多成本。5. 跑通之后如何验证复现质量5.1 计算指标的正确姿势训练完不能只看loss得算Dice系数和IoU。Dice我能直接用def dice_coef(pred, target, threshold0.5): pred (torch.sigmoid(pred) threshold).float() intersection (pred * target).sum() return (2.0 * intersection 1e-6) / (pred.sum() target.sum() 1e-6)注意两个细节第一阈值先固定0.5不要在验证集上反复调否则是变相的过拟合。第二Dice要在原始分辨率上算如果在resize后的224x224上算完再和原图比对边缘区域的误差会被模糊掉。5.2 用少量数据先做sanity check整个复现流程建议分两步走第一步只拿16张图、跑5个epoch。如果训练能正常完成loss从初始值明显下降说明代码链路是通的。这一步还能暴露数据加载、标签维度、设备类型这些低级问题。第二步再上全量数据。正式训练时保存每个epoch的可视化结果把输入图、预测mask、真实mask三张图拼在一起看。loss再漂亮都不如看一眼预测图直观。如果发现某个器官类别完全没预测出来优先检查该类别的像素占比是不是太低了。5.3 后续扩展方向跑通二分类之后你可以往几个方向继续改多类别分割改num_classes和损失函数加入数据增强医学图像里常见的翻转、旋转、弹性形变都能提升鲁棒性替换主干把轻量CNNEncoder换成ResNet50效果更接近原论文加入滑窗推理处理超大尺寸的整张病理切片。如果你想用官方预训练权重复现论文实验核心改动点在PatchEmbed部分要对接ResNet50的输出通道Transformer部分embed_dim改成768depth改成12再把resize_pos_embed用上。轻量版跑通了这些改动都是水到渠成的事。最后说一点我复现完的感受这个模型最值钱的地方不在某一层specific的结构而在于“CNN局部特征 Transformer全局建模 U型跳连”的组合方式。把这条链路弄清楚后面看UNETR、Swin-Unet这些变体时你会发现它们都在同一个框架里做文章。上面这套代码你直接拿去用先跑通二分类再根据自己的数据改路径比我当初摸索时顺畅得多。本文还有配套的精品资源点击获取