
简介面向计算机视觉与深度学习从业者此压缩包提供了TransUnet多分类的完整工程实现覆盖从数据处理、模型搭建到训练评估的整个流程。包内共30个文件包含14个Python源码及对应pyc编译文件可借此快速理解目录结构与运行逻辑同时提供txt训练配置、docx技术说明以及th格式的预训练权重文件整体约373MB。已有5541人学习下载。资源围绕Transformer自注意力机制与U-Net跳跃连接的特点讲解如何将二分类模型改造为多分类结构并详细介绍加权交叉熵、Focal Loss、学习率调度、数据增强等训练策略以及IoU、Precision、Recall、F1等评估指标的使用方法。无论是初学者入门语义分割还是研究者复现与改进模型都能借助其中代码和文档减少环境配置与调参成本特别适用于医学影像分析、遥感图像分类及自动驾驶障碍物识别等场景。1. TransUnet 多分类语义分割从二分类切到多类别时真正要动的几个地方二分类语义分割跑通不难U-Net 最后加一层卷积就能出 mask但任务升级到多分类比如一张眼底图里要同时分出视盘、黄斑和背景三个区域问题立刻具体化输出通道要改成类别总数掩码要以类别 ID 存储而不是 0/1损失函数得换成交叉熵或 Focal Loss评估也不能只看 pixel accuracy。TransUnet 的做法是在 U-Net 编码器末端插入 Transformer 块用自注意力捕获全局上下文再靠跳跃连接把细节特征带回解码器。这套源码工程把训练、推理、评估拆成 train_normal.py、inference.py、eval.py、metrics 与 loss 等模块适合在做医学影像多组织分割、遥感地物分类或自动驾驶场景解析的从业者直接改参数复现。下面按架构、数据、训练、评估四段拆开讲重点放在参数怎么改、改完看什么指标、失败时查哪里。2. TransUnet 架构拆解Transformer 全局注意力与 U-Net 跳跃连接的协同机制2.1 编码器主干与 Transformer 块的接入位置先看 models 目录里的网络定义。常见实现里编码器用 ResNet50 做卷积主干对输入做四次下采样特征图分辨率从 512 降到 32在最后一层特征上把每个像素视为一个 token 送入 Transformer 块做全局自注意力建模。之所以把 Transformer 放在最深层而不是每层都插是因为深层特征感受野已经够大自注意力能高效补上卷积在长程依赖上的短板而浅层插块会显著拖慢训练且收益有限。源码里 TransUNet 的构建参数如下# models/transunet.py 构建示例 import torch from models.transunet import TransUNet model TransUNet( img_dim512, # 输入尺寸须能被 patch_dim 整除 in_channels3, # RGB 三通道 out_channels3, # 多分类类别数背景 两个目标类 head_num4, # 多头自注意力的头数 mlp_dim512, # 前馈网络中间维度 block_num8, # 堆叠的 Transformer 块数量 patch_dim16, # patch 边长决定序列长度 classifierseg # 分割任务头 )逻辑说明patch_dim 与 img_dim 共同决定 Transformer 的序列长度512×512 的图像按 16×16 切分后得到 1024 个 token每个 token 对应原图一个 16×16 区域。自注意力在这 1024 个 token 上做两两交互计算复杂度是序列长度的平方所以 patch_dim 减小到 8 时序列会膨胀到 4096显存直接翻几倍。out_channels 是多分类的关键参数二分类是 2改成 3 就是三分类但它必须和掩码标注的类别数严格一致否则 CrossEntropyLoss 会索引越界。head_num 建议 4 或 8显存吃紧从 4 起步block_num 控制全局建模深度小数据集 6 到 8 足够堆到 12 以上容易过拟合。2.2 跳跃连接与解码器多尺度特征逐级恢复分辨率Transformer 输出的序列要先 reshape 回 32×32 的特征图再进入 U-Net 风格的四级解码器。每一级解码器先做上采样再和编码器对应层的特征做通道拼接之后过两个 3×3 卷积融合。浅层特征保留边缘和纹理细节深层特征提供语义信息这种跨尺度组合是多分类分割里小目标召回率的重要保障。注意自注意力的计算量是序列长度的平方1024 个 token 时单头注意力矩阵大约占 4MBfloat32head_num4 还能接受如果换成 256 输入加 patch_dim8序列变成 4096注意力矩阵膨胀 16 倍很多卡直接 OOM。这也是显存有限的场景优先保 patch_dim16 而不是减 batch 的原因。另外 TransUnet 源码里跳跃连接通常是直接拼接而不是嵌套式密集连接修改 backbone 时编码器各阶段输出通道必须和解码器预期对齐最常见的报错就是这里维度不匹配调试时打印各层 feature shape 就能定位断点。2.3 输出头设计从特征图到类别概率图解码器最后一级输出通道等于 out_channels再接一个 1×1 卷积把特征投影到类别空间得到形状为 [B, C, H, W] 的 logits。训练时 logits 直接送进 CrossEntropyLoss推理时用 argmax 沿类别维取最大索引得到 [B, H, W] 的预测标签图每个像素值就是类别 ID。这里有一个常见误用在输出层先加 Softmax 再做 argmax。Softmax 是单调映射不影响 argmax 结果但会引入额外计算更重要的是如果损失函数换成了 Focal Loss它期望输入是原始 logits提前 Softmax 会让数值分布改变梯度计算对不上。所以 framework.py 里输出头保持裸卷积加 argmax这是多分类分割的标准做法。3. 数据管线与配置dataset 目录、data.py 与 train_normal_config.txt 的配合3.1 数据集目录与掩码格式多分类分割和二分类最大的区别在掩码文件。二分类掩码可以写成黑白图多分类掩码必须是单通道灰度图或索引 PNG像素值直接代表类别 ID——0 是背景、1 是黄斑区、2 是视盘区。项目里 dataset 目录推荐的结构如下dataset/ ├── images/ │ ├── train/ │ └── val/ ├── masks/ │ ├── train/ │ └── val/ └── labels.txtlabels.txt 每行是「类别名 数字 ID」两列脚本按行读取生成 idx_to_label 字典后续可视化、逐类别 IoU 都要靠它把数字还原成名称。掩码如果是从标注工具导出的 RGB 图需要在 data.py 里做一次颜色到 ID 的映射转换做法是建一个颜色表逐像素替换。读取后务必检查类别 ID 集合是否连续。提示训练前在 dataset 的getitem里打印 torch.unique(mask)确认最大类别 ID 等于 num_classes - 1。ID 不连续比如只有 0 和 2会让模型在缺失类别上输出混乱且逐类别 IoU 统计时混淆矩阵形状对不齐。3.2 data.py 的读取与多分类转换下面是 data.py 里适合多分类场景的读取核心# data.py 多分类数据读取核心 import glob import numpy as np import torch from PIL import Image from torch.utils.data import Dataset class SegmentationDataset(Dataset): def __init__(self, image_dir, mask_dir, transformNone): self.image_paths sorted(glob.glob(f{image_dir}/*.png)) self.mask_paths sorted(glob.glob(f{mask_dir}/*.png)) self.transform transform def __getitem__(self, idx): img Image.open(self.image_paths[idx]).convert(RGB) mask Image.open(self.mask_paths[idx]).convert(L) img self.transform(img) mask np.array(mask, dtypenp.int64) mask torch.from_numpy(mask).long() return img, mask说明.convert(L)保证掩码是单通道np.int64再转torch.long()是必须的因为 CrossEntropyLoss 的 target 不接受 Float 型张量这里是最常见的报错点。训练和验证的 transform 要分开写训练用随机翻转加随机旋转加色彩抖动验证只做 resize 和归一化注意旋转要用相同的随机种子同时对 img 和 mask 执行否则标签和图像错位模型会学到错误对应关系。数据增强强度要适度翻转加旋转对小目标密集的场景已经足够过度随机裁剪会切掉目标导致训练信号被稀释。原图和掩码尺寸不一致时确保两者用同一套 resize 参数处理建议在训练前统一到 512×512。3.3 train_normal_config.txt 参数速查config 是这套工程的配置入口训练前应逐项核对参数二分类常见值多分类建议值说明num_classes23N与掩码最大类别 ID 1 一致img_size512512 或 256需能被 patch_dim 整除batch_size848显存不够先减 batch 而非降分辨率loss_typebce_dicece 或 wce多分类不能用 BCElr1e-41e-4AdamW 常用初始值schedulerstepcosine多分类训练建议余弦退火eval_interval55每 N 个 epoch 跑一次验证并保存权重config 里 loss_type 改为 ce 后框架会实例化 loss 目录下的交叉熵实现改成 wce 则使用类别权重版本。验证时打开 save_pred 开关预测图会落盘这是后续排查类别混淆的前提。weights 目录建议按实验命名避免覆盖掉效果较好的 checkpoint。4. 训练主流程train_normal.py 的损失函数、学习率调度与日志记录4.1 多分类损失函数为什么 BCE 不可用Focal Loss 何时需要二分类分割常用 BCE 加 Dice 的组合多分类里 BCE 只输出单通道概率图无法表达多个类别的互斥关系必须换成 CrossEntropyLoss。它的计算方式是先对每个像素的 logits 做 softmax再取目标类别的负对数似然天然支持类别间互斥。当数据存在类别不平衡——比如背景占 90%、目标只占 10%——可以直接在损失里传 class weight 向量权重按类别像素频率的倒数归一化# loss/ce_loss.py 加权交叉熵示例 import torch.nn as nn class WeightedCE(nn.Module): def __init__(self, num_classes, weightNone): super().__init__() # weight: shape [num_classes] 的 Tensor按类别频率倒数设置 self.criterion nn.CrossEntropyLoss(weightweight) def forward(self, logits, target): # logits: [B, C, H, W]target: [B, H, W] 的 Long 张量 return self.criterion(logits, target)逻辑说明CrossEntropyLoss 的 weight 参数会自动完成逐像素加权不需要手动扩展成 [B,1,H,W]。类别权重怎么算最稳统计训练集掩码每个类别的像素数用total_pixels / (num_classes * class_pixels)得到归一化权重再塞进 WeightedCE。如果这样训练后小类别 IoU 还是偏低再考虑 Focal Loss它对难样本分配更大梯度适合目标小且边界模糊的医学影像场景但 Focal Loss 有 alpha 和 gamma 两个超参要调基线不稳时先不要上。4.1.1 类别权重计算示例# 训练前统计类别频率并生成权重 def compute_class_weight(mask_paths, num_classes): counts np.zeros(num_classes, dtypenp.float64) for p in mask_paths: m np.array(Image.open(p).convert(L)) for c in range(num_classes): counts[c] (m c).sum() total counts.sum() return torch.tensor(total / (num_classes * counts), dtypetorch.float32)注意 counts 里有 0 时会出现除零训练前先断言每个类别至少有一个像素否则说明标注文件或类别数配置有问题。4.2 训练循环、学习率调度与断点续训train_normal.py 的训练循环大致如下和 framework.py 里的 trainer 配合# train_normal.py 训练循环核心片段 for epoch in range(start_epoch, epochs): model.train() for imgs, masks in train_loader: imgs, masks imgs.cuda(), masks.cuda() logits model(imgs) loss criterion(logits, masks) optimizer.zero_grad() loss.backward() optimizer.step() if (epoch 1) % eval_interval 0: miou, per_class_iou evaluate(model, val_loader, num_classes) scheduler.step(miou) torch.save({ epoch: epoch, model: model.state_dict(), optimizer: optimizer.state_dict(), }, fweights/epoch{epoch1}.pth)这里有两个细节值得注意。其一scheduler.step() 的调用时机取决于调度器类型余弦退火按 epoch 更新ReduceLROnPlateau 要在验证结束后传入 miou 指标多分类任务建议后者因为 mIoU 比 loss 更能反映收敛状态。其二保存权重时用字典形式把 epoch 和 optimizer 一起存中断训练后可以原地续跑长训练周期里这是必备操作。显存不足时可以加梯度累积每 2 或 4 个 step 累积一次再更新等效放大 batch size同时把学习率线性放大才能对齐收敛行为。4.3 record 与 logs 目录训练过程留痕record 目录存放每轮的损失值和指标logs 里是 tensorboard 事件文件。多分类训练建议至少记录四组曲线总 loss、mIoU、每个类别的 IoU、学习率。只记总 loss 的问题是多分类场景下 loss 下降不一定代表每个类别都在变好——大类别主导 loss 时小类别可能在悄悄变差逐类别 IoU 曲线能第一时间暴露这种状况。记录方式用 SummaryWriter 每轮验证后 add_scalar 写入训练结束后对照曲线判断是继续训练、调权重还是加数据增强。如果某个类别的 IoU 在某个 epoch 后开始下降而整体 loss 还在降说明模型开始牺牲小类别去拟合大类别这时应回到 4.1 调大对应类别的权重。5. 推理与评估inference.py、eval.py 与多分类混淆矩阵的落地写法5.1 推理流程与结果后处理inference.py 加载 weights/ 下选定的 checkpoint对图像做前向推理把 argmax 结果保存为伪彩色图或灰度标签图。保存伪彩色图时要用 labels.txt 的类别映射逐类赋色而不是把整数标签直接当 RGB 输出否则相邻类别的色差太小难以区分。推理时注意关闭梯度并切换到 eval 模式dropout 和 batchnorm 的行为差异在分割任务上会造成明显噪声。5.2 多分类混淆矩阵与 IoU 的 Python 实现评估核心在 eval.py 与 metrics 目录。多分类 IoU 依赖混淆矩阵下面是可直接用的 sklearn 实现以及避免依赖的手写版本# metrics/iou.py 多分类 IoU 与混淆矩阵 from sklearn.metrics import confusion_matrix import numpy as np def compute_iou_per_class(pred, target, num_classes): # pred/target shape 均为 [H, W]像素值为类别 ID cm confusion_matrix( target.flatten(), pred.flatten(), labelslist(range(num_classes)) ) ious [] for c in range(num_classes): tp cm[c, c] fp cm[:, c].sum() - tp fn cm[c, :].sum() - tp denom tp fp fn ious.append(tp / denom if denom 0 else 0.0) return ious, cm def mean_iou(ious): return float(np.mean(ious))代码说明sklearn 的 confusion_matrix 必须传 labels 参数否则某张图里缺失的类别会导致矩阵行数自动压缩IoU 索引随之错位。手写版本等价于逐类别计算 TP / (TP FP FN)分母为零时置 0 而不是报错这是类别缺失时最常见的坑。打印混淆矩阵后对角线是正确分类像素非对角元素大的位置就是最容易混淆的类别对比如类别 1 大量被预测成类别 2就需要回到损失权重或数据增强上找原因。5.3 一个排查技巧全局混淆矩阵定位类别纠缠训练完成后别只盯着 mIoU。把验证集所有预测和标签累计成全局混淆矩阵按行归一化后每一行表示该类别的被误分去向分布。我一般会写个脚本把行归一化矩阵输出成表格或用 matplotlib 存成热力图如果某两个类别互混比例超过 15%先检查 labels.txt 的语义定义和标注边界是否真的可分——标注本身不清晰时模型怎么调都调不好。确认标注没问题后只提高这两个类别的 loss 权重或者给训练数据加局部形变增强通常迭代两个实验就能看到互混比例明显下降。这个先定位再对症下药的闭环比反复堆模型深度更省时间也是 TransUnet 这类混合架构在多分类落地时最值得先做的一步。本文还有配套的精品资源点击获取