简介:一套完整的TransUnet复现工程,面向从事医学图像分割研究或入门Transformer+U-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这些公开数据集上被反复当作baseline,2021年发表在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.0
- torchvision只在需要加载官方预训练权重时用到,基础版本不需要
- 显卡显存建议8GB以上,我用的是RTX 3060 12GB,batch size开到8没问题
- 数据集用标准png格式的图像和mask,图像建议统一到224x224
下面所有代码都直接基于这套环境编写。如果你的PyTorch版本比较老,注意把nn.MultiheadAttention的batch_first=True参数确认一下,老版本不认识这个参数,需要手动把输入维度换成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_ch=3): super().__init__() self.stage1 = nn.Sequential( nn.Conv2d(in_ch, 64, 3, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), ) # 1/2 self.stage2 = nn.Sequential( nn.Conv2d(64, 128, 3, stride=2, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), nn.Conv2d(128, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), ) # 1/4 self.stage3 = nn.Sequential( nn.Conv2d(128, 256, 3, stride=2, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True), nn.Conv2d(256, 256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True), ) # 1/8 self.stage4 = nn.Sequential( nn.Conv2d(256, 512, 3, stride=2, padding=1), nn.BatchNorm2d(512), nn.ReLU(inplace=True), nn.Conv2d(512, 512, 3, padding=1), nn.BatchNorm2d(512), nn.ReLU(inplace=True), ) # 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_ch=512, embed_dim=384, patch_size=1): super().__init__() self.proj = nn.Conv2d(in_ch, embed_dim, kernel_size=patch_size, stride=patch_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_size=1是刻意为之。CNN已经完成了16倍下采样,1/16分辨率特征图上的每个像素就等同于原始输入的16x16 patch,这个写法既保留了TransUNet的patch思想,又省掉了一次reshape的麻烦。
接下来是Transformer部分:
class TransformerBlock(nn.Module): """标准Transformer Encoder Block""" def __init__(self, dim, num_heads, mlp_ratio=4., 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.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_size=2, stride=2) self.conv = nn.Sequential( nn.Conv2d(out_ch + skip_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x, skip): x = self.up(x) x = torch.cat([x, skip], dim=1) return self.conv(x)TransformerBlock里有一个小细节:先把self.norm1(x)存下来再传进attention,而不是在forward里写三次self.norm1(x)。虽然结果等价,但少做两次LayerNorm,前向和反向都会快一点。DecoderBlock的上采样统一用ConvTranspose2d,kernel_size=2、stride=2,刚好把分辨率翻倍,和skip在通道维拼接,再接两层卷积融合。
最后组装成完整的TransUNet:
class TransUNet(nn.Module): def __init__(self, in_ch=3, num_classes=1, patch_size=1, embed_dim=384, depth=6, num_heads=6, dropout=0.1): super().__init__() # CNN编码器 self.encoder = CNNEncoder(in_ch) # 特征图 -> token self.patch_embed = PatchEmbed(in_ch=512, embed_dim=embed_dim, patch_size=patch_size) self.pos_embed = nn.Parameter(torch.zeros(1, (224 // 16) ** 2, embed_dim)) nn.init.trunc_normal_(self.pos_embed, std=0.02) # Transformer编码器 self.blocks = nn.ModuleList([ TransformerBlock(embed_dim, num_heads, dropout=dropout) 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_size=1) 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_size=16对应的196个token。如果你训练时改成512x512,token数变成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, interpolation=cv2.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_ok=True) device = "cuda" if torch.cuda.is_available() else "cpu" dataset = SegmentationDataset("data/images", "data/masks") dataloader = DataLoader(dataset, batch_size=8, shuffle=True, num_workers=4) model = TransUNet(in_ch=3, num_classes=1).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-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(f"Epoch {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_ch=3, num_classes=1).to(device) model.load_state_dict(torch.load("checkpoints/last.pth", map_location=device)) 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。只要输入不是224,196和实际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), mode="bilinear", align_corners=False) pos_embed = pos_embed.permute(0, 2, 3, 1).flatten(1, 2) return pos_embed但我的建议是:除非有明确理由,否则固定输入尺寸。模型是你设计的,没有必要让输入尺寸成为变量。
3.2 损失函数为什么选BCE+Dice
医学图像分割里“背景远多于前景”是常态。以肝脏分割为例,肝脏区域通常只占整个切片的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_workers=4)。Windows环境下如果多进程数据加载报错,可以先改成num_workers=0排查,不是代码逻辑的问题。
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_location="cpu")把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时必须加interpolation=cv2.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", dtype=torch.float16)包住前向和loss计算,显存能省接近一半,训练速度也更快。也可以用torch.utils.checkpoint对Transformer层做梯度检查点,以时间换显存。但最简单的还是把embed_dim从384降到256,depth从6降到4,效果差距不大,显存压力小很多。
4.5 数据加载中的线程和路径坑
补充一个很常见但很容易忽略的问题:Windows下DataLoader(num_workers>0)偶尔会卡死或报BrokenPipeError。这不是代码逻辑问题,而是Windows的spawn方式和Linux的fork不同,建议Windows用户先num_workers=0,等程序完全跑通再尝试调高。数据路径方面,不要在代码里写绝对路径,把data/images这种相对路径和项目根目录绑定,换机器迁移时降低很多成本。
5. 跑通之后,如何验证复现质量
5.1 计算指标的正确姿势
训练完不能只看loss,得算Dice系数和IoU。Dice我能直接用:
def dice_coef(pred, target, threshold=0.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改成768,depth改成12,再把resize_pos_embed用上。轻量版跑通了,这些改动都是水到渠成的事。
最后说一点我复现完的感受:这个模型最值钱的地方不在某一层specific的结构,而在于“CNN局部特征 + Transformer全局建模 + U型跳连”的组合方式。把这条链路弄清楚,后面看UNETR、Swin-Unet这些变体时,你会发现它们都在同一个框架里做文章。上面这套代码你直接拿去用,先跑通二分类,再根据自己的数据改,路径比我当初摸索时顺畅得多。
本文还有配套的精品资源,点击获取