尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

TransUNet:Transformer与CNN融合的医学图像分割实战指南

TransUNet:Transformer与CNN融合的医学图像分割实战指南 1. 从UNet到TransUNet为什么我们需要在医学图像分割中引入Transformer如果你和我一样在计算机视觉领域特别是医学图像分割这个赛道上摸爬滚打过几年那么UNet这个名字对你来说一定像空气一样熟悉。它简洁、高效几乎是所有分割任务的“起手式”。但不知道你有没有遇到过这样的困境面对一张分辨率极高、病灶区域与正常组织边界极其模糊的CT或MRI图像UNet的预测结果总感觉“差那么点意思”——边界不够锐利或者一些微小的、弥散性的病灶区域被漏掉了。这背后的核心原因很大程度上在于UNet的“视野”局限。传统的UNet及其变体如ResUNet、DenseUNet本质上是一个纯卷积神经网络CNN。CNN通过卷积核在局部感受野内提取特征这种归纳偏置局部性、平移不变性是其成功的关键但也成了它的天花板。对于医学图像分割而言一个像素的类别比如是肿瘤还是正常组织往往不仅取决于它周围几个像素的纹理更取决于图像中远距离区域的上下文信息。例如判断一个肺部结节的性质可能需要结合整个肺叶的形态、血管的分布等全局线索。CNN通过堆叠卷积层和下采样来扩大感受野但这个过程是间接且效率较低的信息在传递中容易丢失或模糊难以建立真正长距离的依赖关系。这就是Transformer登场的时候。Transformer最初在自然语言处理NLP领域大放异彩其核心机制“自注意力”Self-Attention能够计算序列中任意两个元素之间的关系权重从而建模全局上下文。TransUNet这篇里程碑式的工作正是将Transformer的这种全局建模能力与UNet的局部特征提取和精确定位能力相结合的一次成功尝试。它不是在取代CNN而是在补足CNN的短板。简单来说TransUNet让模型在分析图像时既能“明察秋毫”CNN负责局部细节又能“纵观全局”Transformer负责上下文关联这对于复杂、多变的医学图像分割任务来说无疑是质的提升。我最初接触TransUNet时也是抱着将信将疑的态度。毕竟Transformer的计算开销和训练难度是出了名的。但当我按照论文思路在自己的数据集上跑通第一个demo并看到其在一些困难样本上显著优于纯CNN模型的分割效果时那种“原来如此”的豁然开朗感让我确信这条路是走对了。接下来我将结合自己的阅读理解和实战训练经验为你拆解TransUNet的核心设计并分享从零开始训练一个TransUNet模型的全过程与避坑指南。2. TransUNet架构深度拆解CNN与Transformer如何“无缝焊接”理解TransUNet关键在于理解它如何将两个看似迥异的架构优雅地融合在一起。整个流程可以清晰地分为三个阶段CNN特征提取、Transformer全局上下文编码、CNN解码上采样。我们一步步来看。2.1 第一阶段CNN骨干网络——提取丰富的局部特征图TransUNet并没有完全抛弃CNN相反它需要一个强大的CNN骨干网络如ResNet-50或ViT的patch embedding层作为“特征提取器”。输入图像例如512x512的CT切片首先通过这个CNN骨干网络。这里有一个关键细节我们并不是取CNN最后的输出而是取其中间层的特征图。为什么是中间层因为CNN的深层特征虽然语义信息丰富知道“这是肿瘤”但空间分辨率太低不知道肿瘤的精确边界。而浅层特征分辨率高、细节丰富但语义性弱。TransUNet通常选取CNN骨干中某个下采样后的特征图例如经过多次下采样后得到的特征图尺寸为原图的1/16或1/32。假设我们输入是(3, 512, 512)经过ResNet-50到某个阶段我们得到一个特征图F其形状为(C, H, W)例如(1024, 32, 32)。这里的C是通道数(H, W)是空间尺寸。这个(1024, 32, 32)的特征图就是我们将要喂给Transformer的“原材料”。它已经包含了由CNN初步加工过的、具有良好局部性的视觉特征。2.2 第二阶段Transformer编码器——建立全局上下文关联这是TransUNet的灵魂所在也是理解上的一个难点。我们需要将二维的图像特征图转换成Transformer能处理的一维序列。步骤1图像序列化Image to Sequence我们将上一步得到的特征图F(1024, 32, 32) 在空间维度上展平。具体操作是把H x W个位置这里是32x321024个位置的每一个都看作是一个“词”。每个“词”是一个C维1024维的特征向量。于是我们得到了一个序列X [x^1, x^2, ..., x^N]其中N H * W 1024每个x^i的维度是C1024。这个序列X的形状是(1024, 1024)即(序列长度N, 特征维度C)。步骤2添加位置编码Positional EncodingTransformer本身不具备感知序列顺序的能力。在NLP中词的位置很重要在图像中像素的空间位置同样至关重要。因此我们必须为序列X中的每一个“词”即每一个图像块的特征向量添加位置信息。TransUNet采用了可学习的位置编码Learnable Positional Encoding即一个与X同形状的可学习参数矩阵P(1024, 1024)。然后执行X X P。这样模型在训练过程中就能学会不同空间位置的重要性。步骤3Transformer编码层Encoder Layers现在这个加上了位置信息的序列X被送入一个标准的Transformer编码器通常是多层堆叠如12层。每一层都包含一个多头自注意力机制Multi-Head Self-Attention, MHSA和一个前馈网络Feed-Forward Network, FNN并且都有残差连接Add和层归一化LayerNorm。多头自注意力机制这是实现全局建模的关键。对于序列中的每一个“词”例如对应原图左上角某个区域的特征自注意力机制会计算它与序列中所有其他词包括距离很远的右下角区域的关联度注意力权重。通过这种机制模型可以学习到“哦这个像素点属于肿瘤不仅因为它周围纹理异常还因为图像另一侧的某个区域出现了典型的卫星灶”。这就是全局上下文。前馈网络对每个位置的特征进行非线性变换和增强。残差连接与层归一化保证训练稳定性和深度网络的信息流动。经过多层Transformer编码后我们得到了一个蕴含了全局上下文信息的特征序列Z其形状仍然是(1024, 1024)。步骤4序列还原Sequence to Feature Map为了后续与CNN解码器对接我们需要把这个一维序列Z还原回二维特征图的形式。因为序列长度N1024对应着原来的空间尺寸H32, W32所以我们直接reshape回去Z(1024, 1024) -Z(1024, 32, 32)。现在我们得到了一个既包含丰富局部细节来自CNN又建模了全局依赖来自Transformer的“增强版”特征图。2.3 第三阶段CNN解码器与跳跃连接——实现精确定位得到了增强特征图Z后TransUNet的后半部分就和一个标准的UNet解码器非常相似了。解码器由多个上采样块组成。每个上采样块通常包含上采样转置卷积或双线性插值 与编码器对应层级特征图的跳跃连接Skip Connection 卷积层。跳跃连接这是UNet的经典设计在TransUNet中至关重要。解码器在上采样过程中会逐级融合来自编码器CNN部分注意不是Transformer输出的对应层级的特征图。这些来自编码器浅层的特征图分辨率高保留了大量的空间细节和边缘信息。通过跳跃连接这些细节被直接“注入”到正在上采样的特征中从而帮助模型恢复出精确的目标边界。上采样与卷积逐步将特征图的空间尺寸放大同时通道数减少最终恢复到输入图像的原尺寸如512x512并输出每个像素的类别概率图。至此TransUNet完成了一次完整的前向传播CNN提取局部特征 - Transformer建立全局关联 - CNN解码器融合局部与全局信息进行精确定位。这个设计巧妙地让两个架构各司其职扬长避短。3. 从零开始训练TransUNet环境、数据与代码实战理论清晰了接下来就是动手环节。训练一个TransUNet模型你需要准备好三样东西环境、数据、代码。我会以在公开医学图像数据集如Synapse多器官分割数据集上训练为例分享我的实操流程。3.1 环境搭建与依赖安装我强烈建议使用conda创建独立的Python环境避免包版本冲突。# 创建并激活环境 conda create -n transunet python3.8 -y conda activate transunet # 安装PyTorch (请根据你的CUDA版本去官网选择对应命令) # 例如对于CUDA 11.3 pip install torch1.12.1cu113 torchvision0.13.1cu113 --extra-index-url https://download.pytorch.org/whl/cu113 # 安装其他核心依赖 pip install numpy opencv-python pillow scikit-learn scikit-image pip install tensorboard # 用于可视化训练过程 pip install einops # 一个非常好用的张量操作库Transformer代码常用 pip install timm # 包含各种预训练CNN骨干网络如ResNet注意PyTorch版本与CUDA版本的匹配是关键。如果版本不匹配会导致无法使用GPU或运行出错。使用nvidia-smi查看CUDA版本然后去PyTorch官网复制对应的安装命令。3.2 数据准备与预处理医学图像数据通常格式特殊如.nii或.dcm且需要对应的标注文件。以Synapse数据集为例它提供了腹部CT的多个器官的3D标注。我们需要将其处理成2D切片用于训练。关键步骤数据读取使用nibabel库读取.nii.gz格式的3D图像和标签。切片提取沿轴向或其他方向将3D体数据切成一系列2D图像。归一化Normalization这是至关重要的一步。CT图像的像素值是亨氏单位HU范围很广如-1000到3000。我们需要将其归一化到一个固定的区间例如[0, 1]。常见做法是采用窗宽窗位Windowing裁剪后再归一化或者对整个数据集的统计量均值和标准差进行归一化。# 示例基于数据集统计的归一化 # 假设已计算好整个训练集的 mean 和 std image (image - mean) / std标签处理分割标签通常是单通道的整数图每个像素值代表类别ID如0背景1肝脏2脾脏...。需要将其转换为one-hot编码形式用于多分类交叉熵损失或保持为长整型用于Dice Loss等。数据增强Data Augmentation医学数据稀缺增强是防止过拟合、提升模型泛化能力的利器。除了常见的旋转、翻转、缩放外医学图像上可以尝试更高级的增强如弹性形变、伽马变换、对比度调整等。albumentations库是一个很好的选择。数据集划分按照病人ID划分训练集、验证集和测试集切记不能按随机切片划分否则会导致来自同一病人的数据同时出现在训练和测试中造成数据泄露使评估结果虚高。3.3 模型构建代码核心解析网上有很多TransUNet的开源实现。选择一个结构清晰、易于修改的代码库至关重要。以下我提炼几个关键部分的代码逻辑1. Transformer编码器块import torch.nn as nn from einops import rearrange class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio4., qkv_biasFalse, drop_rate0.): super().__init__() self.norm1 nn.LayerNorm(dim) self.attn nn.MultiheadAttention(dim, num_heads, dropoutdrop_rate, biasqkv_bias) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(drop_rate), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(drop_rate) ) self.drop_path nn.Dropout(drop_rate) if drop_rate 0. else nn.Identity() def forward(self, x): # x: (序列长度N, 批次大小B, 特征维度C) # 自注意力残差连接 x x self.drop_path(self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]) # 前馈网络残差连接 x x self.drop_path(self.mlp(self.norm2(x))) return x2. 将CNN特征图转换为序列def embed_features(cnn_feature_map): cnn_feature_map: (B, C, H, W) 输出: (N, B, C) 其中 N H * W B, C, H, W cnn_feature_map.shape # 将空间维度展平并调整维度顺序以适应Transformer x cnn_feature_map.flatten(2).transpose(1, 2) # (B, N, C) x x.transpose(0, 1) # (N, B, C) return x3. 损失函数选择医学图像分割中常用的损失函数是Dice Loss和交叉熵损失Cross-Entropy Loss的加权和。因为医学目标往往占比较小类别不平衡Dice Loss直接优化分割区域的重叠度对此类问题很有效。class DiceCELoss(nn.Module): def __init__(self, weight_ce1.0, weight_dice1.0): super().__init__() self.weight_ce weight_ce self.weight_dice weight_dice self.ce nn.CrossEntropyLoss() def dice_loss(self, pred, target): # pred: (B, C, H, W) after softmax # target: (B, H, W) LongTensor smooth 1e-6 target_one_hot F.one_hot(target, num_classespred.shape[1]).permute(0, 3, 1, 2).float() intersection (pred * target_one_hot).sum(dim(2,3)) union pred.sum(dim(2,3)) target_one_hot.sum(dim(2,3)) dice (2. * intersection smooth) / (union smooth) return 1 - dice.mean() # 平均各类别的Dice Loss def forward(self, pred, target): ce_loss self.ce(pred, target) pred_softmax F.softmax(pred, dim1) dice_loss self.dice_loss(pred_softmax, target) total_loss self.weight_ce * ce_loss self.weight_dice * dice_loss return total_loss4. 训练过程中的核心技巧与避坑指南有了代码和数据训练过程才是真正考验功力的地方。以下是我在多次训练TransUNet中积累的经验和踩过的坑。4.1 学习率策略与优化器选择Transformer模型通常对优化策略比较敏感。我推荐使用AdamW优化器它是Adam的改进版解耦了权重衰减通常能获得更好的泛化性能。学习率设置是关键。一个常见的策略是使用带热启动Warmup的余弦退火Cosine Annealing学习率调度器。Warmup在训练初期例如前5%的步数或轮数将学习率从一个很小的值如1e-7线性增加到初始学习率如1e-4。这有助于稳定训练初期防止梯度爆炸。Cosine Annealing在Warmup之后按照余弦函数将学习率从初始值衰减到接近0。这比阶梯式下降更平滑往往能找到更优的解。from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR from torch.optim import AdamW optimizer AdamW(model.parameters(), lrbase_lr, weight_decay1e-4) # 先定义warmup scheduler warmup_scheduler LinearLR(optimizer, start_factor0.01, total_iterswarmup_epochs * steps_per_epoch) # 再定义cosine annealing scheduler 从warmup结束开始 cosine_scheduler CosineAnnealingLR(optimizer, T_max(total_epochs - warmup_epochs) * steps_per_epoch) # 实际训练循环中先执行warmup_scheduler.step()再执行cosine_scheduler.step()4.2 解决显存溢出OOM问题TransUNet尤其是当输入图像较大、Transformer层数较多时显存消耗非常恐怖。自注意力机制的计算复杂度是序列长度的平方O(N²)我们的序列长度N1024这已经不小了。应对策略减小批次大小Batch Size这是最直接的方法但可能会影响BN层的统计和训练稳定性。可以考虑使用梯度累积Gradient Accumulation。例如你想用批次大小16但显存只够4那么你可以设置实际批次大小为4但累积4步后再更新一次梯度optimizer.step()等效于批次大小16。accumulation_steps 4 for i, (images, labels) in enumerate(train_loader): loss model(images, labels) loss loss / accumulation_steps # 损失归一化 loss.backward() if (i1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()降低输入分辨率这是最有效的方法之一。如果任务允许尝试将输入图像从512x512降到256x256序列长度N会从65536降到4096显存和计算量会大幅下降。需要权衡精度损失。使用混合精度训练AMP使用torch.cuda.amp自动将部分计算转换为半精度float16可以显著节省显存并加速训练。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()检查点梯度Gradient Checkpointing这是一种用时间换空间的技术只保存部分中间结果在反向传播时重新计算其余部分。对于非常深的Transformer模型很有效。PyTorch中可以通过torch.utils.checkpoint.checkpoint实现。4.3 模型评估与指标解读训练时不能只看损失函数下降必须在独立的验证集上监控分割指标。医学图像分割常用的指标有Dice相似系数Dice Coefficient最核心的指标衡量预测区域与真实区域的重叠度。Dice 2 * |A ∩ B| / (|A| |B|)。值越接近1越好。通常报告各类别的平均DicemDice。豪斯多夫距离Hausdorff Distance, HD衡量两个轮廓之间的最大不匹配程度对边界分割精度非常敏感。值越小越好。由于对异常点敏感常用95%分位数HD95。交并比IoU与Dice类似IoU |A ∩ B| / |A ∪ B|。在验证时要同时观察这些指标。有时损失在下降但Dice不升反降可能是过拟合的迹象。要保存验证集上指标最好的模型而不是训练损失最低的模型。4.4 一个常见的“坑”位置编码与输入尺寸这是我在复现时踩过的一个大坑。可学习的位置编码P的形状是固定的(N, C)其中N H * W。这意味着如果你在训练时使用的输入图像尺寸是(512, 512)经过CNN下采样后特征图尺寸是(32, 32)那么N 1024。你的位置编码P就是(1024, C)。问题来了如果你在测试或部署时想处理一个不同尺寸的图像比如(640, 480)经过同样的CNN下采样后特征图空间尺寸变了假设变成(20, 15)N300。此时你预训练的位置编码P(1024, C) 就无法直接与新的序列 (300, C) 相加了解决方案固定输入尺寸在训练和推理时使用完全相同的输入尺寸。这是最简单的方法但缺乏灵活性。使用插值在模型加载预训练权重后对位置编码P进行二维插值使其匹配新的空间尺寸。这需要将一维序列P先reshape回(H, W, C)然后进行插值再展平。这种方法有一定效果但并非最优因为位置编码的语义可能被插值破坏。使用相对位置编码或条件位置编码一些改进的Transformer变体如Swin Transformer中的相对位置偏置或CPVT中的条件位置编码能更好地处理可变尺寸输入。但这需要对模型结构进行修改。因此在项目开始前务必确定好你的输入尺寸策略。对于医学图像通常可以统一重采样到固定尺寸这是最稳妥的做法。训练一个TransUNet模型是一场耐心和细节的较量。从数据预处理的一个参数到损失函数的一个权重再到学习率调度器的一个周期都可能对最终结果产生显著影响。我的经验是严格按照论文描述设置基线然后在一个小规模验证集上进行快速的超参数扫描如学习率、权重衰减、损失函数权重找到适合你自己数据的最优配置再开始全量训练。这个过程没有捷径但每一次成功的训练都会让你对模型和数据有更深的理解。
返回列表