TransUnet这个网络,常跑医学图像分割的朋友应该都不陌生。它把CNN的特征提取能力和Transformer的全局建模能力拼在一起,在不少分割任务上都拿到了不错的效果,现在很多论文还是会拿它当对比基准。但这里有个很现实的问题:官方代码默认是跑Synapse、ACDC这类医学灰度影像的,数据加载部分直接读的是npy或者h5格式,预处理链路也完全围绕灰度图设计。如果你手里只有一堆彩色RGB图片,直接拿官方train.py开跑,大概率会在数据读取阶段就报错,或者训练出来的模型根本学不到东西。这篇文章把我自己的改造过程完整拆一遍,从数据集整理、官方代码结构解读、RGB三通道适配,到训练参数调整和推理落地,一条龙讲清楚,帮想用TransUnet跑自己RGB分割数据的朋友少走弯路。
1. 先弄清TransUnet做了什么,以及你的RGB数据为什么不能直接喂给它
1.1 网络结构核心思路
TransUnet的结构可以简单拆成三块:CNN编码器、Transformer编码器、U-Net解码器。CNN部分通常用ResNet50作为骨干,先对输入图像做下采样提取局部特征;随后把CNN输出的特征图切成固定大小的patch,展平成序列后送入Transformer的Encoder里,用自注意力机制去捕获像素之间的长距离依赖关系;最后再把Transformer输出的序列重新还原成特征图的形状,交给U-Net风格的解码器逐步上采样,恢复出和原图分辨率一致的分割结果。
这种混合设计的核心理由其实很直白:纯Transformer擅长建模全局关系,但对像素级细节有点粗枝大叶,直接用在分割上容易出现边界糊成一片的问题;纯CNN又受限于感受野,很难把相隔很远的同类区域连起来。TransUnet相当于用CNN先打底、用Transformer补全局视野,再用U-Net把细节捞回来,三个环节各司其职。实际训练中,被称作“三明治”式的混合编码器带来的效果改善在医学影像这种结构复杂、边界模糊的任务上尤其明显。
1.2 官方代码默认的数据流与RGB数据的冲突点
TransUnet官方仓库的数据处理流程,核心是围绕Synapse多器官分割数据集设计的。官方数据加载模块会把原始影像和标签都转成npy格式的数组,然后按2D切片的方式一张张喂给网络。这个流程里默认的假设是输入是灰度图:数据增强、归一化、通道维度的处理全按单通道来。而Synapse这类数据集,标签图里的每个像素值对应一个器官类别编号,整体是索引标注图而不是彩色标注图。
RGB三通道彩色图替换进来之后,主要会撞上三堵墙:
- 数据加载模块读图方式不对。官方代码很多地方直接假设输入已经是npy数组,而不是从JPG、PNG这类常见图片格式里读RGB值。
- 归一化参数不对。灰度图的归一化一般只针对单通道做,而RGB图像需要逐通道做标准化,否则模型输入分布完全偏离预训练权重的分布,收敛速度会慢很多。
- 标签处理逻辑不对。很多自己标的数据集,标注文件保存成彩色PNG,每个类别用一种颜色;这类彩色标注必须转成0、1、2...这样的类别索引图,才能配合交叉熵这类损失函数使用。
所以这里有一个很关键的结论:TransUnet的网络结构本身是能吃RGB三通道输入的,ViT部分也有对应的3通道patch embedding,真正需要改的是官方代码里“从文件到网络输入”之间的这一整条数据链路。这篇文章后面讲的所有改动,本质上都是围绕这条链路来做手术。
2. 训练环境准备与数据集整理规范
2.1 硬件环境和依赖安装
我自己用的是一张RTX 4090,24GB显存,batch size可以开到16配合224x224输入尺寸,跑150个epoch大约需要六到八个小时。如果你手里的显卡是12GB显存,建议batch size先降到8,或者用梯度累积来模拟更大的batch。CPU内存方面,32GB比较稳妥,因为数据增强阶段会在内存里同时解压多张图处理,遇到大佬数据集动辄上万张图时内存小了很容易被系统杀掉进程。
环境依赖上,官方代码仓库要求的核心依赖是:
python 3.8+ pytorch 1.10+ torchvision timm einops tensorboard SimpleITK(跑Synapse原版数据时需要,跑纯RGB数据可暂时不装)装完基础依赖后,我额外建议安装opencv-python和albumentations,前者读取图片方便、速度快,后者做图像增强时能保持image和mask同步变换,比手动写增强函数省心很多。安装命令如下:
pip install opencv-python albumentations tensorboard einops timm2.2 数据集目录结构设计
整理自己的RGB分割数据集,我强烈建议一开始就按下面的目录规范来组织,后面改代码、写加载器、做评估都会省很多事:
datasets/ ├── RGBDataset/ │ ├── train/ │ │ ├── images/ # 训练原图,jpg/png均可,3通道 │ │ │ ├── img_001.jpg │ │ │ └── img_002.jpg │ │ └── masks/ # 训练标签,PNG格式,单通道索引图或彩色标注图 │ │ ├── img_001.png │ │ └── img_002.png │ ├── val/ │ │ ├── images/ │ │ └── masks/ │ ├── test/ │ │ ├── images/ │ │ └── masks/ │ └── train.txt # 每一行:图片路径 制表符 标签路径 └── pretrained/ ├── R50+ViT-B_16.npz # TransUnet官方提供的ImageNet预训练权重 └── R50+ViT-B_16.json关于标签的格式,这里有个重点需要单独拎出来说。我见过太多人在这个环节掉坑:用Labelme或者PS标注完之后,保存出来的mask是彩色图,比如类别0是黑色、类别1是红色、类别2是绿色,每个类别对应一种RGB向量。这种情况下不能直接把maks图喂给网络,因为网络输出的类别索引是0、1、2这种一维数字,和RGB颜色向量对不上。必须先把彩色标注图转成单通道索引图,做法是把每个RGB值映射到一个类id上。
如果标注图是单通道PNG,里面每个像素的灰度值恰好等于类别id,那就省事了,直接读就行。在Python里用opencv读图时注意一点:cv2.imread(path, cv2.IMREAD_GRAYSCALE)读出来的是单通道灰度图,cv2.imread(path)默认读出来的是BGR三通道彩色图,这两个模式千万别搞混。
2.3 数据划分与类别映射表
划分数据集时,我建议按**训练集70%、验证集15%、测试集15%**的比例随机划分。随机划分前先检查一下每张图的类别分布,确保验证集和测试集里每个类别都存在,不然评估阶段算出来的mIoU / Dice会忽高忽低,没有参考意义。更精细一点的做法是按图像来源分组划分,比如某一批图来自同一台设备或者同一时段拍摄,就把它们放到同一组里再按组切分,防止数据泄漏导致的指标虚高。
手动标注出来的RGB分割数据集,类别数量从二分类到十来个类别都很常见。通常需要维护一张类别映射表,类似这样:
| 类别id | 像素颜色RGB | 含义 |
|---|---|---|
| 0 | (0, 0, 0) | 背景 |
| 1 | (255, 0, 0) | 目标物A |
| 2 | (0, 255, 0) | 目标物B |
| 3 | (0, 0, 255) | 目标物C |
写转换脚本时,优先把所有像素RGB值用numpy的矩阵运算一次性映射到类别id,不要用Python循环逐像素判断,后者在1920x1080这种分辨率下会慢到怀疑人生。代码示例如下:
import numpy as np import cv2 color_to_id = { (0, 0, 0): 0, (255, 0, 0): 1, (0, 255, 0): 2, (0, 0, 255): 3, } def rgb_mask_to_index(mask_path, out_path): # 读取彩色标注图,注意OpenCV默认是BGR顺序 mask_bgr = cv2.imread(mask_path) mask_rgb = cv2.cvtColor(mask_bgr, cv2.COLOR_BGR2RGB) h, w = mask_rgb.shape[:2] # 初始化索引图,-1表示未知类别,方便后面检查漏标 index_map = np.full((h, w), fill_value=-1, dtype=np.int32) for cls_id, rgb in color_to_id.items(): match = np.all(mask_rgb == np.array(rgb).reshape(1, 1, 3), axis=-1) index_map[match] = cls_id if np.any(index_map == -1): print(f"警告: {mask_path} 中存在未归类像素") # 保存为单通道PNG cv2.imwrite(out_path, index_map.astype(np.uint8))3. 核心代码改造:把官方数据链路换成RGB版本
3.1 官方代码结构导读
改代码之前,先把官方仓库的文件布局摸清楚。TransUnet官方仓库里,与训练直接相关的关键文件大概有这几个:
train.py # 训练入口 test.py # 测试入口 datasets/ ├── dataset_synapse.py # Synapse数据集加载器 └── dataset_acdc.py # ACDC数据集加载器 lib/models/transunet.py # TransUnet模型定义 utils/ ├── losses.py # 损失函数 ├── metrics.py # 评估指标 └── test.py # 测试工具函数 lists/ ├── lists_Synapse/ │ ├── train.txt │ └── test_vol.txt其中dataset_synapse.py是我花时间最多的地方,因为官方代码里大量使用了np.load直接读npy数组,并内置了“把3D体数据切成2D切片”的逻辑。我们自己跑RGB图片,最省事的方案其实不是去官方代码上修修补补,而是直接写一个新的Dataset类,把官方数据集类和文件路径都替换掉。
3.2 自定义RGB Dataset类
下面这份代码是我在实际项目中整理出来的,可以直接放到datasets/目录下新建一个dataset_rgb.py文件里:
import os import cv2 import numpy as np import torch from torch.utils.data import Dataset import albumentations as A from albumentations.pytorch import ToTensorV2 class RGBSegmentationDataset(Dataset): def __init__(self, data_dir, split_file, img_size=224, mode='train', num_classes=2): self.data_dir = data_dir self.mode = mode self.img_size = img_size self.num_classes = num_classes # 读取split文件,每行是"images/xxx.jpg\tmasks/xxx.png" self.samples = [] with open(split_file, 'r', encoding='utf-8') as f: for line in f: line = line.strip() if not line: continue img_rel, mask_rel = line.split('\t') img_path = os.path.join(data_dir, img_rel) mask_path = os.path.join(data_dir, mask_rel) if os.path.exists(img_path) and os.path.exists(mask_path): self.samples.append((img_path, mask_path)) # 数据增强管线 if mode == 'train': self.transform = A.Compose([ A.RandomResizedCrop(height=img_size, width=img_size, scale=(0.7, 1.0), p=0.5), A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.2), A.RandomBrightnessContrast(p=0.3), A.Normalize( mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], max_pixel_value=255.0 ), ToTensorV2() ]) else: self.transform = A.Compose([ A.Resize(height=img_size, width=img_size), A.Normalize( mean=[123.675, 116.28, 103.53], std=[58.395, 57.12, 57.375], max_pixel_value=255.0 ), ToTensorV2() ]) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, mask_path = self.samples[idx] # 读取RGB原图:cv2默认读成BGR,先转回RGB image_bgr = cv2.imread(img_path) if image_bgr is None: raise ValueError(f"无法读取图像: {img_path}") image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) # 读取标签:如果标签是单通道索引图,用IMREAD_GRAYSCALE mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) if mask is None: raise ValueError(f"无法读取标签: {mask_path}") # 标签里如果带有未知类别(255),需要先处理掉 mask = np.clip(mask, 0, self.num_classes - 1) # 数据增强:image和mask必须保持相同的空间变换 augmented = self.transform(image=image_rgb, mask=mask) image_tensor = augmented['image'] mask_tensor = torch.from_numpy(augmented['mask']).long() return image_tensor, mask_tensor这段代码里有几个地方是专门为RGB三通道数据设计的。归一化用的mean和std就是ImageNet预训练权重对应的RGB通道数值,顺序是R、G、B。如果你用的是自己从零训练的模型,也可以改用自己数据集上统计的均值和方差,但大多数情况下用ImageNet的数值已经足够。
此外,我用albumentations而不是PyTorch官方的torchvision.transforms做增强,原因是这个库自带了一套image和mask同步变换的封装。拿随机裁剪来说,如果只用torchvision自带的RandomCrop,原图和标签如果不用同一个随机种子去裁剪,位置会对不上;albumentations里直接传两个参数进去,它内部帮你处理好同步问题,代码简洁、也不容易出bug。这一点在分割任务里非常实用。
3.3 修改训练入口文件
官方train.py里与数据加载和模型输出相关的几个关键改法,我逐个说明。
第一个是实例化Dataset和DataLoader的部分。官方代码里创建dataloader的语句可能长这样:
db_train = Dataset(base_dir=args.root_path, list_dir=args.list_dir, split="train") trainloader = torch.utils.data.DataLoader(db_train, batch_size=args.batch_size, shuffle=True, num_workers=8, pin_memory=True)这行代码里的Dataset指的是官方自带的Synapse_dataset。我们要替换成自己写的RGBSegmentationDataset,并且传入split文件路径而不是目录:
from datasets.dataset_rgb import RGBSegmentationDataset train_dataset = RGBSegmentationDataset( data_dir=args.root_path, split_file=args.list_dir, # 这里改成train.txt文件路径 img_size=args.img_size, mode='train', num_classes=args.num_classes ) trainloader = torch.utils.data.DataLoader( train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=8, pin_memory=True, drop_last=True )这里drop_last=True是个小细节。如果训练集样本数不能被batch_size整除,丢弃最后一个不完整的batch,能避免BatchNorm层在训练和推理时行为不一致的问题。样本总数不足时这个设置尤其重要。
第二个是num_classes参数要传对。官方代码里经常在训练脚本顶部硬编码了输出类别数,比如n_classes=9对应Synapse的8个器官加背景。改成自己的数据集时,务必把这里替换成自己数据集的类别总数,包括背景类。否则模型最后一层输出的通道数和标签中的最大类别索引对不上,训练时损失直接爆掉,报错信息多是“Target size != Tensor size”之类。
第三个是模型定义部分。TransUnet模型在构造时通常需要传入img_size和in_chans等参数:
net = TransUnet( img_size=224, in_chans=3, num_classes=args.num_classes, embed_dim=768, depth=12, num_heads=12, ... )in_chans=3是RGB三通道的关键,千万别改成1。虽然官方仓库的预训练模型通常是在3通道ImageNet上预训练的,但这个参数还是要显式写清楚。如果你的环境里加载模型时因为timm库版本问题报错,可以考虑升级或降级timm到0.4.12左右,官方代码依赖的是比较老的版本。
3.4 预训练权重的加载处理
TransUnet通常会在Vit部分使用ImageNet上预训练的R50+ViT-B_16权重。加载预训练权重时,有一个容易翻车的点:官方提供的R50+ViT-B_16.npz文件解压出来的权重dict里,键名和模型当前状态的键名可能不一致,直接load_state_dict会报“Missing key(s) and unexpected key(s)”的警告。
我的习惯是写一个简单的兼容加载函数:
import numpy as np import torch def load_pretrained_npz(model, npz_path): npz = np.load(npz_path, allow_pickle=False) weights = {k: torch.from_numpy(v) for k, v in npz.items()} model_dict = model.state_dict() matched = {} for k, v in weights.items(): if k in model_dict and model_dict[k].shape == v.shape: matched[k] = v model_dict.update(matched) model.load_state_dict(model_dict) print(f"成功加载预训练权重: {len(matched)} / {len(model_dict)} 层匹配") return model这种加载策略允许网络部分结构不匹配时仍然能加载成功,不匹配的层保留随机初始化,后面训练时这些层会自己学出来。实测下来,用预训练权重做初始化,比从零开始训练在收敛速度和最终Dice指标上普遍要好5到10个百分点,所以有条件的话尽量把这个权重用上。
4. 训练参数配置与训练流程实操
4.1 优化器、学习率与损失函数的选择
官方代码默认用的是SGD优化器加CosineAnnealing学习率调度。我自己在RGB分割任务上,更推荐AdamW + CosineAnnealing + 线性warmup的组合。SGD在医学分割这种小数据集上收敛偏慢且对学习率敏感,AdamW对learning rate没那么挑剔,配合warmup预热可以让Transformer部分在训练初期更稳定。
我常用的一组基础参数如下:
优化器: AdamW 初始学习率: 1e-4 权重衰减: 1e-4 batch size: 16(24GB显存)/ 8(12GB显存) 训练轮数: 150 学习率调度: CosineAnnealing,50轮warmup结束到1e-3再降回1e-5如果你执意用SGD,学习率可以设置为0.01到0.05之间,配合poly学习率策略,效果也不错。不过从我测试来看,AdamW在分类不平衡的数据集上表现更稳定,尤其当你的数据里背景像素占比很大的时候。
损失函数部分,官方提供了DiceLoss和CrossEntropyLoss两个选择,我强烈建议两者结合起来一起用:
ce_loss = nn.CrossEntropyLoss() dice_loss = DiceLoss(num_classes=args.num_classes) total_loss = ce_loss(logits, labels) + dice_loss(logits, labels)单独用CrossEntropyLoss时,如果背景像素占90%以上,网络很容易把所有像素都预测成背景,Dice会很难看;单独用DiceLoss时,又容易出现loss震荡不收敛的情况。两个loss加起来,CE负责提供稳定的梯度方向,Dice负责对齐类间平衡,配合起来效果最稳。权重比例我一般设1:1,如果你的类别极度不均衡,可以加大DiceLoss的权重,比如0.6的系数给dice、0.4给CE。
4.2 训练过程中的监控指标解读
训练时我习惯同时打印这几个指标:train loss、train dice、val dice、val mIoU。loss能反映模型是否在收敛,dice能反映分割质量和背景类占比是否合理。这里说一个经常遇到的“假收敛”现象:如果你的分类问题里背景像素占绝大多数,模型很快就能学会全部预测成背景,这时候train loss看起来在下降,但val dice可能只有0.3甚至更低。应对办法就是上面说的DiceLoss+CrossEntropy组合,以及在训练中打印按类别分别统计的dice,看清楚模型到底是哪个类别学不动。
我在训练过程中的tensorboard日志通常会记录以下几类曲线:
train/loss train/dice val/dice val/iou val/loss lr其中val/iou里的IoU,也就是交并比,是衡量分割结果和真实标注重合度的最直观指标。计算方法是:每个类别分别算预测正确的像素数 / (预测为该类的像素数 + 真实为该类的像素数 - 预测正确的像素数),最后对所有类别取平均得到mIoU。如果某个类别在训练集里只有几十个像素,它的IoU对整体均值影响不大,但你一旦做测试集评估,这个类别几乎必然预测失败,因此要在训练之前就想好要不要对这样的rare class做上采样或者类别加权。
4.3 显存不够时的节流三板斧
实测中,即使24GB显存,如果把分辨率调到512x512,batch size基本只能开到4到6。如果显存继续爆掉,按下面顺序检查:
第一,把batch size调到2甚至1试试,这是最简单直接的办法。batch size小会让BatchNorm的统计量不稳,如果你的batch size小于8,建议把网络里的BatchNorm换成GroupNorm或者InstanceNorm,不然验证时使用running mean,指标会有明显跳动。
第二,开混合精度训练。PyTorch的自动混合精度(AMP)可以把显存占用砍掉将近一半,而且对最终精度基本没有影响。核心改动只有几行代码:
scaler = torch.cuda.amp.GradScaler() for batch_idx, (images, labels) in enumerate(trainloader): images, labels = images.cuda(), labels.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): logits = net(images) loss = ce_loss(logits, labels) + dice_loss(logits, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()第三,把num_workers调低一些。有时候爆显存不是显存真的不够,而是DataLoader线程数太多导致内存线程切换开销大,num_workers=4到8基本够用了,不用盲目开满。
4.4 训练中断恢复与checkpoint管理
训练跑到一半因为服务器重启、显存被占用等原因挂掉,是家常便饭。我的习惯是每5个epoch保存一个checkpoint,文件名里带上epoch号和val dice,并且单独保存一份最新的“last.pth”,方便随时从断点恢复。
torch.save({ 'epoch': epoch, 'model_state_dict': net.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'best_dice': best_dice, }, 'checkpoints/last.pth')恢复训练时只需要重新加载last.pth,把网络、优化器、调度器的状态全部恢复,再继续跑就行。如果数据增强管线里用了random seed,记得把np.random.seed和torch.manual_seed也一并固定,否则恢复之后虽然模型状态一样,但增强出来的数据完全不同,训练曲线可能出现小跳变。
5. 验证与推理:从训练好的模型到一张张预测图
5.1 测试指标的计算与可视化
训练结束后,在测试集上计算Dice、IoU这些指标时,要特别注意一点:TransUnet输出的是一个形状为B, num_classes, H, W的概率图,需要对每个像素在类别维度上做argmax,得到最终的类别索引图。然后拿这个索引图和真实标签逐像素比较。
def calculate_metrics(pred_mask, true_mask, num_classes): iou_list = [] dice_list = [] for cls in range(num_classes): pred_cls = (pred_mask == cls) true_cls = (true_mask == cls) intersection = (pred_cls & true_cls).sum() union = (pred_cls | true_cls).sum() if union == 0: iou_list.append(float('nan')) # 该类别在GT中不存在 else: iou_list.append(intersection / union if union > 0 else 0.0) dice_list.append(2 * intersection / (pred_cls.sum() + true_cls.sum()) if (pred_cls.sum() + true_cls.sum()) > 0 else 0.0) return np.nanmean(iou_list), np.nanmean(dice_list)遇到某个类别在真实标签中完全不存在的测试图,iou和dice会碰到除零问题。直接跳过该类别还是给0分,取决于你评估的目的。如果是给论文写实验对比,我一般建议按“该类在ground truth中出现才参与计算”的规则处理,并在题注里注明;如果是实际产品验证,则更倾向于给0分,因为漏检的情况同样需要被惩罚。
5.2 单张图片推理演示
实际部署时经常需要对任意一张新图做在线预测。下面这份脚本可以直接在命令行里跑,加载训练好的权重,输入一张RGB图片,输出预测mask和彩色叠加图:
import torch import cv2 import numpy as np from lib.models.transunet import TransUnet def predict_image(model, image_path, img_size=224, device='cuda'): model.eval() image_bgr = cv2.imread(image_path) image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) original_h, original_w = image_rgb.shape[:2] # resize到网络输入尺寸 image_resized = cv2.resize(image_rgb, (img_size, img_size), interpolation=cv2.INTER_LINEAR) # 归一化 mean = np.array([123.675, 116.28, 103.53], dtype=np.float32) std = np.array([58.395, 57.12, 57.375], dtype=np.float32) image_norm = (image_resized.astype(np.float32) - mean) / std image_tensor = torch.from_numpy(image_norm.transpose(2, 0, 1)).unsqueeze(0).float().to(device) with torch.no_grad(): logits = model(image_tensor) # (1, num_classes, img_size, img_size) pred = torch.argmax(logits, dim=1).squeeze(0).cpu().numpy() # (img_size, img_size) 类别索引 # 恢复到原图分辨率 pred_resized = cv2.resize(pred.astype(np.uint8), (original_w, original_h), interpolation=cv2.INTER_NEAREST) return pred_resized def save_overlay(image_path, pred_mask, output_path): image_bgr = cv2.imread(image_path) overlay = image_bgr.copy() # 这里定义简单颜色映射,可根据自己的类别调整 color_map = { 1: (0, 0, 255), # 红:目标A 2: (0, 255, 0), # 绿:目标B 3: (255, 0, 0), # 蓝:目标C } for cls, color in color_map.items(): overlay[pred_mask == cls] = color blended = cv2.addWeighted(image_bgr, 0.5, overlay, 0.5, 0) cv2.imwrite(output_path, blended)这个脚本里有个关键细节:resize预测结果时插值方式必须用INTER_NEAREST,不能使用INTER_LINEAR或INTER_CUBIC。因为分割mask是离散的类别编号,用线性插值会在类别边界产生不存在的中间值,比如类别1和类别2之间出现0.6这种小数值,转回uint8后会成为另一个错误类别,导致边缘出现一圈鬼影。同理,保存在推理阶段所有涉及预测mask的空间变换,都要用最近邻插值。
5.3 预测结果的后处理与常见问题
做完预测后,你大概率会看到两类问题。一类是细小孔洞,就是预测结果里本该连续的大块区域中间出现零星的小洞。这种情况推荐用简单的形态学闭运算填掉,常用的是cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel),kernel大小设为3x3或5x5就够,太大会把细长结构也给抹掉。
另一类是类别边界毛刺感很强。这是像素级分割的通病,可以用CRF后处理技术改善,但引入CRF会显著增加推理时间,实际工程里要权衡。如果你的场景只在乎分割的大致区域,对边缘精细度要求不高,那就保留原始输出即可,不用额外后处理。
6. 常见问题与排错实战记录
6.1 读图报错与通道顺序错误
我自己在最初跑RGB数据时,踩的第一个坑就是OpenCV的通道顺序。OpenCV的cv2.imread返回的数组是BGR顺序,而PyTorch模型训练时一般期望RGB顺序。如果不做cv2.COLOR_BGR2RGB转换,就相当于把红通道和蓝通道对调后送进网络,模型训练时loss可能照样下降,但验证效果奇差无比,因为模型学到的颜色语义完全是错位的。排查方法很简单:把读出来的图存一张可视化对比一下,看看红色物体是否显示成了蓝色。
6.2 标签类别索引与输出类别数不匹配
如果标签图中最大像素值是4,但num_classes设置成了3,训练时会在计算损失时报错:Target 4 is out of bounds。相反,如果num_classes设置大了而标签里最大类别只有2,倒是能跑,但模型会浪费大量参数去学一个不存在的类别。所以训练前务必要检查一下标签的最大值:
mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) print("label min:", mask.min(), "label max:", mask.max())一旦发现max值等于255,几乎可以确定标签是彩色标注图被当成灰度图读进来了,需要先做颜色到类别的映射。
6.3 验证集Dice高但测试集Dice低
这种情况大多不是过拟合,而是训练集和测试集分布不一致。RGB分割里最常见的原因是图像的光照条件不同,比如训练数据是从晴天场景采集的,测试集里全是阴天或者夜晚图像。此时仅靠数据增强是不够的,需要针对性做色彩增强,比如在训练管线中加入随机色调扰动、灰度化、对比度抖动等。
一个更隐蔽的问题是图像分辨率不一致。如果训练时统一resize到了224x224,而测试图像是1920x1080,再resize回224x224,小目标细节信息会损失严重。我的建议是,训练时就模拟测试阶段的分辨率策略,如果测试时会把整张大图切块预测再拼回来,那么训练时也应该用相同尺寸的切块去做增强和裁剪。
6.4 训练loss不下降或直接变NaN
训练loss变成NaN,最常见的原因是学习率设置过大,梯度更新一步就越过了数值稳定边界。这时候把学习率降到1e-5到1e-4量级再试,情况基本能缓解。第二个常见原因是标签中有超出num_classes范围的异常值。第三个原因是从npy或h5文件里读到了包含NaN的输入数据,比如原图某些像素出现坏值,训练时梯度回传就会出问题。
如果loss一直不下降,先别急着调参,把训练数据可视化一遍。我之前遇到过数据文件夹里混入了大量损坏的图片,OpenCV读取时返回None,但代码里没有做检查直接硬算,导致loss波动非常离谱。跑训练之前,批量扫描一遍图片能否正常读取是个很好的习惯:
import cv2, os from tqdm import tqdm bad_files = [] for root, dirs, files in os.walk('datasets/RGBDataset'): for f in files: if f.lower().endswith(('.jpg', '.png', '.jpeg')): path = os.path.join(root, f) img = cv2.imread(path) if img is None: bad_files.append(path) print(f"损坏图片数量: {len(bad_files)}")6.5 数据加载慢、GPU利用率上不去
训练时GPU利用率只有30%左右,CPU反而跑满,这是典型的数据加载瓶颈。虽然我们把num_workers开到8了,但如果每个worker处理单张图时都要做大量CPU操作,比如大尺寸resize、彩色mask转换、复杂增强,整体速度依然会被拖慢。几个行之有效的优化手段:
- 图片尺寸在进入DataLoader前先统一缩放,避免每个worker都在做接近原始分辨率的resize操作。
- 用
pin_memory=True让GPU直接访问锁页内存,减少数据拷贝时间。 - 数据增强管线里避免使用逐像素的Python循环,尽量用OpenCV和numpy的向量化操作。
- 如果训练集规模本来就几万张,可以考虑在离线阶段把所有图预处理并缓存成LMDB或TFRecord格式,训练时直接读取缓存。这个方案改动较大,一般数据量小的时候不需要上。
6.6 跑通了但分割效果一言难尽怎么办
如果训练正常跑完,预测时mask却是一团噪点或者大片漏分割,大概率不是代码问题,而是模型能力没有充分释放。我建议按固定顺序排查:
先确认评估指标本身可靠。打印出测试集前10张图的预测mask,转成灰度图可视化,人眼判断预测结果是否至少有大致形状。如果只是边缘粗糙,说明模型学到了结构信息,可以靠后处理或者继续训练提升。
再核查数据集内标注质量。有些标注的边界很粗糙,类别和原图对不齐,这会导致Dice上限被拉低,再怎么调参也上不去。把原图和标注半透明叠加显示,检查几处边界是否贴合得干净。
最后再考虑模型本身的变化,比如增大输入分辨率、增加训练轮数、换更强的预训练权重。TransUnet在224x224输入下,对细长结构的分割天然吃亏,因为Transformer的patch尺寸是16x16,一条只有几个像素宽的线在patch内可能直接糊化。有条件的话可以把输入改成384x384或512x512。显存不够时就用上面的混合精度和梯度累积方案腾空间。
7. 写在最后的几点实操建议
如果只让我给一个最核心的忠告,那就是:第一次跑通全流程,永远先用最小样本集。我当时先用20张图、10个epoch跑了一遍完整流程,确认loss能下降、checkpoint能保存、推理脚本能出图,才动用全量数据开正式训练。这一步看着多花了半小时,实际上帮你节省的是排查“训练跑了一半才发现mask读错了”这种灾难现场的时间。
另外,TransUnet官方代码本身是研究性质的代码,训练脚本写得相对粗糙,很多地方没有做错误处理。改造时不要有“官方的一定是完美”的执念,按自己数据的特点去改数据加载器和训练策略,才是正路。这篇文章里给的Dataset类和推理脚本是我在RGB分割任务上跑过多次的基础版本,你可以直接拿去改路径跑通,再根据具体业务调整增强策略、后处理逻辑和类别权重。
RGB三通道图像的分割,核心工作量从来不在网络结构,而是在数据链路。把这层窗户纸捅破,后面自然顺畅。