ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

TransUnet多分类医学图像分割实战:从标签编码到评估全流程

2026/9/11 22:38:59 拓冰建站 浏览量
TransUnet多分类医学图像分割实战:从标签编码到评估全流程 简介一份面向深度学习语义分割算法研究与实践的TransUnet多分类代码包聚焦Transformer全局注意力与U-Net局部特征提取的混合架构适用于医学影像分析、遥感图像分类、自动驾驶场景理解等任务。内容覆盖模型搭建、数据读取、训练、推理与评估全流程并包含加权交叉熵/Focal Loss等损失函数实现帮助学习者快速掌握从二分类拓展到多类别分割的模型改造与训练技巧。资源共30个文件包括14个Python源码、13个编译后的pyc文件、训练配置txt、方法说明docx及模型权重文件压缩包约373MB。随附说明文档对数据增强、学习率调度、正则化等策略进行了梳理代码中给出IoU、Precision、Recall等评估指标接口目录按训练、推理、数据、权重等模块组织并带有日志记录便于复盘实验可直接运行和二次开发。已有5541人学习下载尤其适合具备一定深度学习基础、希望结合Transformer与CNN落地多类分割方案的研究者和开发者。1. TransUnet多分类从二分类思维跳到像素级标签矩阵做过分割的人都知道TransUnet 在医学影像里火起来是因为它用 CNN 提取局部特征、用 Transformer 建模全局依赖正好补上了 U-Net 感受野不足的短板。但网上能找到的案例十有八九是二分类背景一类、目标一类。真到了多分类场景——比如把肝脏 CT 里的肝脏、肿瘤、血管、脾脏同时分出来——很多人把代码里的num_classes改成 5 就以为完事了结果训练时 loss 直接 NaN或者预测出来整张图都是背景。多分类和二分类的差别不在网络结构而在标签编码、损失函数、评估指标这三个层面。标签从单通道 0/1 变成单通道 0/1/2/3/4或者变成 one-hot 的五通道损失从 BCE 变成 CrossEntropyLoss评估从 IoU 变成 mIoU 和类别混淆矩阵。这三个点没改对网络再先进也白搭。这篇文章就沿着「数据准备 → 模型改参 → 训练配置 → 评估验证」这条线把 TransUnet 多分类落地时的每一个关键开关讲清楚。适合已经跑通过 U-Net 二分类、现在要扩到多类别的工程师也适合第一次接触 TransUnet 但想绕过常见坑的人。2. 数据与标签多分类的起点是 mask 编码方式不是通道数2.1 单通道标签与 one-hot 标签的区别TransUnet 的输入是(B, C, H, W)的图像C 是输入通道数通常为 1灰度或 3RGB。但标签的形态决定了你后面用哪条 loss 分支。常见做法是训练集里的 mask 是单通道 PNG每个像素值直接是类别编号背景为 0第一类为 1第二类为 2依此类推。比如肝脏分割数据集中像素值 0 是背景1 是肝脏2 是肿瘤。这种编码方式叫稀疏标签也是绝大多数医学影像数据集的原始格式。另一种是 one-hot 标签形状是(B, num_classes, H, W)每个类别一个通道该类别所在的通道为 1其余为 0。one-hot 的好处是方便可视化每个类别的概率输出但会显著增加内存和磁盘占用。一个 512×512×5 的 one-hot mask 是 5 个通道而单通道稀疏标签只有 1 个通道。TransUnet 官方实现里默认使用的是CrossEntropyLoss这个 loss 在 PyTorch 中要求输入是原始的类别索引张量即稀疏标签。所以如果你把 mask 转成了 one-hot再用 CrossEntropyLoss就得先torch.argmax(output, dim1)还原成索引否则会报错或计算错误。我一般建议数据集原始是稀疏标签就保持稀疏标签不要人为转 one-hot。只有在使用 DiceLoss 并实现多分类版本时才需要把稀疏标签转成 one-hot因为 Dice 系数的计算要对每个类别单独做。# 读取多分类 mask 并检查类别分布 import numpy as np import cv2 mask cv2.imread(label.png, cv2.IMREAD_UNCHANGED) print(mask shape:, mask.shape) print(unique values:, np.unique(mask)) # 输出示例 # mask shape: (512, 512) # unique values: [0 1 2 3 4]这段代码告诉你两件事mask 必须是单通道否则np.unique会返回三维数组类别数要从 0 开始连续编号不能出现 0、1、3 这种跳跃否则CrossEntropyLoss的类别索引会错乱。2.2 TransUnet 的类别数参数容易被忽略的位置开源社区里流传最广的 TransUnet 实现入口参数里有几个和类别数直接相关的字段。改num_classes是最表面的一步真正容易漏的是下面几个位置。第一如果用了官方的预训练权重encoder部分的out_channels不用改但decoder最后的卷积层输出通道必须等于num_classes。第二模型内部的分类 token 或嵌入维度如果不匹配会在 forward 阶段直接报 shape 错误。第三数据增强时如果用了随机裁剪mask 的插值方式不能用默认的双线性否则类别边界会出现 0.2、0.7 这种非整数像素值。# 以典型的 TransUnet 初始化为例 model TransUNet( img_dim512, in_channels3, num_classes5, # 这里必须是类别总数包含背景 head_num4, mlp_dim512, block_num8, patch_dim16, classifierseg )注意num_classes5意味着你要分出背景 四个目标类别。很多人把目标类别数当成num_classes导致预测时 logits 通道数比预期少 1后续 mIoU 计算全是错的。2.3 类别不平衡多分类的 mask 往往长这样多分类分割里最典型的问题是类别极度不平衡。背景像素可能占 95%某个小目标类别只占 1%。直接用交叉熵模型会倾向于把所有像素预测为背景因为这样 loss 最低。一个实用的做法是给每个类别设置权重。权重可以基于像素频率的倒数也可以根据先验知识手工设置。TransUnet 训练时在 loss 里传weight参数即可。import torch import torch.nn as nn # 假设统计出每类像素占比背景0.9, 类别10.05, 类别20.03, 类别30.015, 类别40.005 class_weights torch.tensor([0.2, 1.0, 2.0, 4.0, 8.0], dtypetorch.float32).cuda() criterion nn.CrossEntropyLoss(weightclass_weights)权重的设置逻辑占比越小的类别权重越大。具体数值不是固定的我一般先跑一次完整 epoch统计各类别的像素比例再按median_freq / class_freq调整。median_freq是所有类别频率的中位数这样不会让某个极小类别的权重爆炸。如果你发现加权重之后小类别 recall 上来了但 precision 掉了可以把权重的指数调小比如从[0.2, 1, 2, 4, 8]改成[0.5, 1, 1.5, 2, 3]。这是一个经验性的平衡过程没有公式可以直接套。# 统计数据集中每个类别的像素数量 def compute_class_freq(mask_list): freq np.zeros(5) for path in mask_list: m cv2.imread(path, cv2.IMREAD_UNCHANGED) for cls_id in range(5): freq[cls_id] np.sum(m cls_id) return freq / freq.sum()这段统计代码建议放在训练脚本的最前面输出一次频率表并记录下来方便写报告时引用。很多实验对比不是模型差异而是类别权重没对齐。3. 模型改参与 forward 输出从 U-Net 到 TransUnet 多分类的 4 个硬性调整3.1 输入尺寸与 patch 大小的匹配关系TransUnet 的核心机制是把输入图像切成固定大小的 patch然后序列化送入 Transformer。patch 大小决定了序列长度序列长度直接影响注意力矩阵的显存占用和计算量。假设输入是512×512patch 大小是 16那么序列长度是(512/16)² 1024。如果输入是256×256patch 还是 16序列长度变成256。序列长度减到四分之一注意力矩阵计算量直接减少到十六分之一显存压力大幅下降。实际使用中我通常在 2D 医学影数据集上选img_dim256或224patch_dim16这样既能保留足够的空间分辨率又不会因为序列太长导致显存溢出。如果你只有单张 12GB 显存的卡输入 512 且 batch size 为 4 时不加梯度累积大概率会 OOM。此时可以把输入缩到 256或者把 batch size 降到 1再用梯度累积模拟更大 batch。# 训练时观察显卡显存占用 nvidia-smi --query-gpumemory.used,memory.total --formatcsv -l 2显存不够时不要直接砍网络深度先缩输入尺寸。因为 TransUnet 的 Transformer 部分对序列长度的敏感性远高于 CNN 部分缩输入带来的影响比砍 block 数小得多。3.2 decoder 输出通道必须等于 num_classes但别忘上采样TransUnet 的解码器部分通常包含多层上采样和卷积。最后一层卷积的输出通道应等于num_classes但输出的空间尺寸不一定等于原始输入尺寸。如果你在模型外部接上采样或者在 loss 计算时直接对 label 做 resize都能工作但前者更常见。# 假设模型 forward 返回的是 (B, num_classes, H, W) logits model(image_batch) print(logits.shape) # 期望 torch.Size([B, 5, H, W]) # 如果 logits 的空间尺寸和 label 不一致用插值对齐 from torch.nn.functional import interpolate if logits.shape[-2:] ! label.shape[-2:]: logits interpolate(logits, sizelabel.shape[-2:], modebilinear, align_cornersFalse)这段代码解决的是多分类中最隐蔽的一个问题数据增强时的随机裁剪会让 label 和 image 尺寸一致但模型内部的 stride 或 padding 可能导致输出特征图尺寸略小于输入。如果你确定模型是 encoder-decoder 结构且和输入尺寸相同可以不加 interpolation但保留这一段更安全。3.3 分类头在 TransUnet 里不参与分割输出TransUnet 同时支持分割和分类两种任务。在classifierseg模式下模型返回的是分割 logits在classifiercls模式下返回的是图像级分类结果。多分类分割要确保传给 loss 的是 seg 输出而不是把两个输出拼接在一起。有些实现里模型会返回一个 tuple包含分割输出和辅助分类输出。如果你在训练循环里直接写loss criterion(output, label)而这个output是 tuplePyTorch 会直接报类型错误。保险起见在训练脚本里显式取第一个元素。output model(image_batch) if isinstance(output, (tuple, list)): output output[0] loss criterion(output, label_batch)这段防御性代码值得保留因为在不同版本的 TransUnet 实现中返回值结构可能不一致。3.4 冻结 encoder 权重加速多分类微调如果你用的是在 ImageNet 或自然图像上预训练的权重迁移到医学影像多分类任务时可以选择冻结 encoder 的前几层。Transformer 的前几层学的是通用局部模式比如边缘、纹理冻结后可以减少参数量更新降低过拟合风险同时加快训练。# 冻结模型 encoder 部分参数 for name, param in model.named_parameters(): if encoder in name or name.startswith(vit): param.requires_grad False注意 decoder 部分的参数不冻结因为上采样和最后的分类卷积需要针对你的数据集重新训练。如果你发现 loss 降得很慢可以把requires_grad改为 True再继续训练几个 epoch。冻结策略是微调期才用的不是从头训练时的常规手段。4. 训练策略与损失函数多分类的 loss 不只是把 BCE 换成 CE4.1 CrossEntropyLoss 的多分类本质是 N 个二分类的求和PyTorch 的CrossEntropyLoss内部会先对 logits 做 softmax然后取目标类别对应的负对数概率。它把多分类问题拆成相互独立的二分类问题每个类别有自己的概率输出。这里有一个容易误解的地方softmax 输出的所有类别概率之和等于 1这意味着类别之间是互斥的。如果你的分割目标里某些像素可能同时属于两个类别——比如一个区域既是肿瘤也是血管——那 softmax 交叉熵就不合适。此时要换成 multi-label 的 BCE 或带 sigmoid 的 DiceLoss。TransUnet 多分类最主流的选择依然是CrossEntropyLoss因为医学影像的解剖结构类别通常是互斥的。下面是一个完整的多分类训练循环片段包含 loss 计算、反向传播和梯度裁剪。import torch.optim as optim optimizer optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100) criterion nn.CrossEntropyLoss(weightclass_weights) for epoch in range(epochs): model.train() for images, labels in train_loader: images images.cuda() labels labels.cuda() optimizer.zero_grad() logits model(images) if isinstance(logits, (tuple, list)): logits logits[0] loss criterion(logits, labels) loss.backward() # 梯度裁剪防止 Transformer 部分梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step()clip_grad_norm这行不是可有可无的。TransUnet 的 Transformer encoder 部分在训练初期梯度范数可能非常大尤其是在 batch size 较小时。加上梯度裁剪后训练稳定性会明显提升不会出现 loss 突然跳到 NaN 然后永久无法恢复的情况。4.2 多分类 DiceLoss 的正确打开方式DiceLoss 对类别不平衡的鲁棒性通常优于 CrossEntropyLoss但多分类 DiceLoss 的实现是一个很容易写错的地方。常见的错误是先把整个 batch 的所有像素拉平成一个一维向量然后计算 Dice这样会把不同类别的空间信息混在一起导致结果没有意义。正确的做法是对每个类别单独计算 Dice然后取平均。def multiclass_dice_loss(logits, labels, num_classes5, smooth1e-5): probs torch.softmax(logits, dim1) labels_onehot torch.nn.functional.one_hot(labels, num_classesnum_classes) # labels_onehot shape: (B, H, W, num_classes) - (B, num_classes, H, W) labels_onehot labels_onehot.permute(0, 3, 1, 2).float() dice 0.0 for cls in range(num_classes): pred probs[:, cls] target labels_onehot[:, cls] intersection (pred * target).sum() union pred.sum() target.sum() dice (2.0 * intersection smooth) / (union smooth) return 1.0 - dice / num_classes这段代码里有个关键点one_hot操作必须在 GPU 上完成吗不一定但permute和float转换不能省略。你可以在 CPU 上做 one_hot再.cuda()但为了减少数据拷贝我倾向直接在本来的设备上操作。4.3 CELoss 与 DiceLoss 的加权混合实际项目中只用 CrossEntropyLoss 往往对小目标不够敏感只用 DiceLoss 在训练初期会非常不稳定。常见的做法是把两个 loss 加权相加权重可以固定也可以按 epoch 动态调整。import torch.nn.functional as F ce_weight 0.5 dice_weight 1.0 def combined_loss(logits, labels): ce F.cross_entropy(logits, labels, weightclass_weights) dice multiclass_dice_loss(logits, labels, num_classes5) return ce_weight * ce dice_weight * dice我一般从ce_weight0.5, dice_weight1.0开始如果验证集 mIoU 在小目标类别上偏低就把 dice 权重加大到 1.5 或 2.0。注意两个 loss 的量级不一致CE 通常在 1~3 之间Dice 在 0.2~0.8 之间所以直接用相等的权重让 dice 占据主导是合理的。4.4 常用超参速查表参数建议值说明img_dim224 或 256显存 8GB 以下选 224patch_dim16不宜大于 32否则细节损失batch_size4~8配合梯度累积使用learning_rate1e-4 ~ 3e-4Transformer 部分用较小 lr 更稳weight_decay1e-5 ~ 1e-4防止过拟合epochs80~150多分类建议用 early stopping这个表是通用起点不是金科玉律。你的数据集如果非常小比如只有几十张图lr 降到 5e-5epochs 加到 200配合强数据增强更稳。5. 实战流程一个多分类分割训练脚本的完整骨架5.1 数据集封装与预处理多分类任务的数据加载器必须在__getitem__里同时读取图像和 mask并且保证两者的数据增强方式一致。最稳妥的方式是使用同一个随机种子对 image 和 label 做同样的几何变换。import torch from torch.utils.data import Dataset import cv2 import albumentations as A class SegmentationDataset(Dataset): def __init__(self, image_paths, mask_paths, img_size256): self.image_paths image_paths self.mask_paths mask_paths self.img_size img_size self.aug A.Compose([ A.RandomCrop(img_size, img_size), A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(p0.2), ]) def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image cv2.imread(self.image_paths[idx]) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_UNCHANGED) # 用相同的 transform 处理 image 和 mask augmented self.aug(imageimage, maskmask) image augmented[image] mask augmented[mask] image torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0 mask torch.from_numpy(mask).long() return image, mask这里有个容易踩的坑mask必须转成torch.long()否则CrossEntropyLoss会报错因为 target 需要是整数且类别索引在有效范围内。另外RandomCrop和HorizontalFlip都是几何变换对 mask 是安全的但RandomBrightnessContrast只作用于 image不能写在 mask 上albumentations 的Compose默认知道哪些变换支持 mask哪些不支持所以这里不需要手动区分。数据加载器里的collate_fn在 batch 内图像的 shape 必须一致。如果数据集中原始图像尺寸不同RandomCrop会统一尺寸所以在上面的实现里不会出问题。如果你不做随机裁剪就得在__getitem__里用Resize统一尺寸。self.aug A.Compose([ A.Resize(img_size, img_size), A.HorizontalFlip(p0.5), A.RandomBrightnessContrast(p0.2), ])5.2 训练主循环与 checkpoint 保存多分类任务里只保存 loss 最低的模型是不够的因为 loss 低不代表小类别的分割效果好。我习惯同时保存 mIoU 最高的模型并在每个 epoch 结束后做一次验证。def validate(model, val_loader, criterion): model.eval() total_loss 0.0 with torch.no_grad(): for images, labels in val_loader: images images.cuda() labels labels.cuda() logits model(images) if isinstance(logits, (tuple, list)): logits logits[0] loss criterion(logits, labels) total_loss loss.item() return total_loss / max(len(val_loader), 1)验证集不参与反向传播torch.no_grad()可以省显存同时加速推断。注意model.eval()会把dropout和batch_norm切换到推理模式这会影响结果。如果你发现训练集 mIoU 高但验证集 mIoU 特别低先检查是不是忘了切eval()然后再怀疑过拟合。保存 checkpoint 时不仅要保存模型权重还要保存优化器状态和当前 epoch。因为多分类训练通常要跑上百个 epoch中间可能中断有了完整 checkpoint 可以无缝恢复。torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_miou: best_miou, }, checkpoint_epoch%d.pth % epoch)恢复训练的代码不复杂但要确保模型结构参数和保存时完全一致包括num_classes。如果你改过类别数再从头加载旧权重最后一层 shape 不匹配会直接报错。此时只加载不包含最后一层的权重即可或者干脆从头训。5.3 训练日志里应该看哪些关键指标多分类训练过程中除了整体 loss还要分别看每个类别的 IoU。整体 mIoU 可能因为背景类别占比大而看起来很高比如 0.95但实际上肿瘤类别的 IoU 只有 0.2。我习惯在验证阶段输出一个字典包含每类的 IoU 和整体 mIoU。下面的代码用了 Python 的多分类混淆矩阵思想先统计每个类别的 TP、FP、FN再逐个计算 IoU。def compute_per_class_iou(pred, label, num_classes5): ious [] for cls in range(num_classes): pred_mask (pred cls) label_mask (label cls) intersection (pred_mask label_mask).sum().item() union (pred_mask | label_mask).sum().item() iou intersection / union if union 0 else 0.0 ious.append(iou) return ious输出结果形如[0.98, 0.85, 0.62, 0.31, 0.48]。如果某个类别的 IoU 为 0说明模型完全没预测出这个类别。原因可能是样本太少、权重不够大、或者数据增强把这类样本的边缘破坏得太厉害。此时优先调 loss 权重而不是改网络结构。每个 epoch 的验证输出建议写成 CSV 或 JSON放在日志目录里。这样后期可以画每个类别的 IoU 曲线对比不同 epoch 的变化趋势找过拟合的时间点。6. 验证技巧与 Python 多分类混淆矩阵在 TransUnet 评估中的应用评估 TransUnet 多分类分割模型最直观的工具就是混淆矩阵。sklearn.metrics.confusion_matrix可以直接接受两个一维数组分别表示真实标签和预测标签但分割任务的输出是二维图。你需要把每个像素的类别预测和真值都拉平成一个列表再传入函数。import numpy as np from sklearn.metrics import confusion_matrix, classification_report all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in val_loader: images images.cuda() logits model(images) if isinstance(logits, (tuple, list)): logits logits[0] preds torch.argmax(logits, dim1).cpu().numpy() all_preds.extend(preds.flatten()) all_labels.extend(labels.numpy().flatten()) cm confusion_matrix(all_labels, all_preds, labels[0, 1, 2, 3, 4]) print(cm)这段代码里的labels参数很重要。因为某些类别可能完全没有出现在验证集的某张图里如果不显式指定混淆矩阵的维度会少于类别数后续可视化会错位。labels列表的顺序就是最终矩阵的行列顺序和num_classes保持一致。除了混淆矩阵classification_report可以输出每个类别的 precision、recall、f1-score。分割任务里这几个指标按像素统计能直接反映模型对每个类别的敏感度。比如类别 1 的 recall 是 0.7说明有 30% 的真实类别 1 像素被预测成了背景或其他类别。可视化混淆矩阵时用matplotlib的imshow配合colorbar和xticklabels来标记类别名。颜色越深代表数值越大对角线上的值越高越好。import matplotlib.pyplot as plt def plot_confusion_matrix(cm, class_names): fig, ax plt.subplots(figsize(8, 8)) im ax.imshow(cm, interpolationnearest, cmapplt.cm.Blues) ax.figure.colorbar(im, axax) ax.set(xticksnp.arange(cm.shape[1]), yticksnp.arange(cm.shape[0]), xticklabelsclass_names, yticklabelsclass_names, xlabelPredicted label, ylabelTrue label) plt.setp(ax.get_xticklabels(), rotation45, haright, rotation_modeanchor) plt.show()对 TransUnet 多分类模型来说混淆矩阵能一眼看出哪些类别容易被混淆。比如肝脏和肿瘤的像素在边缘区域经常重叠肿瘤区域小但形状不规则漏检率通常较高。如果混淆矩阵里类别 2 有很大一部分被分到了类别 1说明这两个类别的纹理特征过于相似Transformer 的全局注意力可能反而弱化了局部细节此时可以适当增加更高分辨率的 CNN 特征或者在 loss 里对这组易混淆类别单独加权。最后补充一个实际技巧预测结果保存时不要直接存 logits 或概率图而是存argmax后的单通道 PNG。PNG 是无损压缩适合作为最终分割结果交付。同时保存一张彩色编码的可视化图把不同类别映射成不同的 RGB 颜色便于人工检查。用 OpenCV 的applyColorMap是其中一种便捷做法但多类别时更推荐自己构建一个固定的颜色列表确保每次输出的同一类别颜色一致方便横向对比不同 epoch 的结果。对 TransUnet 这类参数规模比较大的模型评估阶段也建议使用torch.cuda.amp.autocast()混合精度推理可以显著降低显存占用并提高吞吐量同时不会明显影响精度。多分类验证集较大时这个优化值得加到评估脚本里。本文还有配套的精品资源点击获取