十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

用Res2Net改进UNet实现舌头图像语义分割实战

用Res2Net改进UNet实现舌头图像语义分割实战 简介这套基于UNet与Res2Net模块改进的舌头图像语义分割项目以PyTorch为框架面向医学影像分析、深度学习入门及语义分割进阶的研究者与学生提供从数据预处理到训练评估的完整流程。资源包共610个文件约7.37MB包括300张JPG原图、300张PNG掩码图以及4份Python代码、1份项目说明书DOCX和说明文档TXT/MD数据集与代码一一对应。项目支持二分类与多类别分割整合数据增强、自动标签处理、IoU/Dice评估并可通过命令行配置数据路径、学习率与标签映射训练后输出模型权重、曲线和指标日志便于直接使用或二次改进。目前已有59人学习下载适合需要快速构建分割基线或深入研究Res2Net改进策略的读者。1. 舌头图像语义分割为什么要动 UNet 的结构把舌头图像分割这件事做扎实难点从来不是“跑通一个 UNet”而是舌头本身太不配合舌尖和舌根的色差大舌苔边界是渐变的裂纹和齿痕又细又浅普通 UNet 在编码器下采样时会把小结构直接丢掉最后出来的掩码边缘经常是“碎”的。Res2Net 模块恰好能在不加深网络的前提下把同一层的感受野宽度拉开让编码器在保留全局轮廓的同时不牺牲细粒度纹理这正是舌头分割最需要的特性。这篇文章不是给你一份现成项目说明书而是顺着“UNet Res2Net 模块改造 舌头数据集 完整代码”这条线把每一层的设计意图、参数选择和落地坑位讲清楚。适合已经跑通过 UNet、想在医学或细粒度语义分割任务上做改进的工程师和研究生看完你至少能回答三个问题Res2Net 放在 UNet 哪个位置收益最大、舌头数据集怎么标怎么增强不翻车、训练时哪个超参数对结果影响排在第一位。2. Res2Net 原理拆解以及它凭什么改进 UNet2.1 Res2Net 的多尺度粒度和普通空洞卷积不是一回事Res2Net 发表在 CVPR 2020它的核心改动非常小在残差块的内部把经过1x1卷积压缩后的特征图按通道维度切成s份一般s4从第二份开始每一份都会经过一个3x3卷积且输入是前一份的输出加上当前份的特征。这样从第二份往后每一条分支的等效感受野是逐级放大的网络在同一层里就拥有了从3x3到3x3*s的连续尺度覆盖。这和 ASPP、空洞卷积系列的区别在于ASPP 是在特征图外侧并联不同 dilation rate权重是共享的Res2Net 是串联式的逐步融合更接近“特征金字塔”在单层内的微缩版计算量增量却小得多。对舌头分割而言舌体轮廓是大尺度目标舌裂、齿痕是中尺度舌乳头纹理是小尺度三类特征如果不能在同一层同时出现解码器后期就很难融合出干净的边界。用公式表达一个Res2Block的前向过程假设输入经过1x1压缩后得到x按通道切成s份x_ii1,2,...s定义y_i为第i份的输出y_1 x_1 y_2 conv3x3(x_2 y_1) y_3 conv3x3(x_3 y_2) y_s conv3x3(x_s y_{s-1})实际工程中可以做两种变体一种让y_1也过一次3x3所有分支都统一另一种是y_1直接跳连如上式。两者的 mIoU 差距在 0.5% 以内但后者省一次卷积训练更快。我的默认选项是后者。2.2 UNet 中放置 Res2Net 的三个候选位置对比把 Res2Net 模块塞进 UNet位置选择会影响最终指标的 3~5 个百分点这不是玄学是不同深度对多尺度特征的需求强度不同。常见做法有三种先看对比表插入位置做法收益点副作用编码器全部卷积块把每层两个3x3卷积替换为 Res2Block各层同步获得多尺度最稳参数量增加约 20%显存占用上升仅最底层瓶颈层第 4 层替换为 Res2Block语义信息最丰富收益高浅层细节仍然丢失跳跃连接处在 skip connection 前加一个 Res2Block融合浅层细节多尺度化对深层的全局尺度无能为力我在舌头数据集上的实际体验是三选二组合编码器前两层 瓶颈层性价比最高。第一层和第二层分辨率高感受野小细粒度纹理主要靠这两层保留瓶颈层控制全局语义。如果全部替换显存占用上涨而第三层本身是中等语义替换后对结果的提升和它带来的训练时间不成比例。所以接下来给出的完整代码采用“前两层 Res2Block 瓶颈层 Res2Block”的改进方案第三层保持普通卷积解码器不动这也是这套改进在小型医学数据集上收敛最快、最不容易过拟合的配置。2.3 面向 UNet 改造的 Res2Block PyTorch 实现下面这个 PyTorch 实现直接可用不依赖任何第三方外部库只基于torch.nn。这里用的是BasicBlock结构适配 UNet 每层的通道数变化。import torch import torch.nn as nn class Res2Block(nn.Module): def __init__(self, in_channels, out_channels, scale4, stride1): super().__init__() self.scale scale # 1x1 降维控制计算量width 是每个分支的通道数 width out_channels // scale self.conv1 nn.Conv2d(in_channels, width * scale, kernel_size1) self.bn1 nn.BatchNorm2d(width * scale) # 中间的多尺度 3x3 卷积第一分支不参与因为 y1 x1 self.convs nn.ModuleList([ nn.Conv2d(width, width, kernel_size3, padding1, stridestride) for _ in range(scale - 1) ]) self.bns nn.ModuleList([ nn.BatchNorm2d(width) for _ in range(scale - 1) ]) self.conv3 nn.Conv2d(width * scale, out_channels, kernel_size1) self.bn3 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) def forward(self, x): identity x out self.relu(self.bn1(self.conv1(x))) xs torch.chunk(out, self.scale, dim1) ys [] fuse xs[0] for i in range(self.scale - 1): if i 0: fuse xs[0] fuse fuse xs[i 1] y self.relu(self.bns[i](self.convs[i](fuse))) ys.append(y) fuse y # 第一分支原样保留拼接后 1x1 恢复通道 ys [xs[0]] ys out torch.cat(ys, dim1) out self.bn3(self.conv3(out)) if identity.shape out.shape: out identity return self.relu(out)参数说明scale4时参数量约为普通两个3x3卷积的 1.2 倍显存增加约 15%20%stride参数是给下采样层用的stride2时在3x3上直接降采样能省掉一层池化。代码里的torch.chunk是按通道切分切分维度和scale必须整除如果out_channels是 64scale设为 4每个分支 16 个通道这个取值在 UNet 第一层表现不错。3. 舌头数据集制作从标注到可训练的完整代码3.1 标注类别的选择直接决定网络学习难度舌头分割不是“舌头一个类、背景一个类”这么简单。实际做中医舌诊辅助系统时至少要把舌头区域拆成两类舌体不含舌苔的舌质部分和舌苔。这两类的边界在很多样本里是渐变的只标一个前景类会让网络在渐变带上产生严重的不确定预测。类别设定建议0背景1舌质2舌苔。如果你的任务更细比如还要分割齿痕或裂纹单独开类会导致样本不均衡更推荐先做二分类前景分割再在 ROI 内部做细分类的两阶段方案。一阶段直接分 4 类以上在几百张数据上几乎必然收敛困难。我通常用 Labelme 标注每张图生成一个 JSON 文件记录多边形顶点。但 Labelme 原生的 JSON 转掩码方式速度慢且不容易做多类合并所以我自己写了转换逻辑。3.2 使用 labelme 半自动标注后的 JSON 转掩码脚本import json import base64 import numpy as np import cv2 import os def labelme_json_to_mask(json_path, shape(512, 512)): with open(json_path, r, encodingutf-8) as f: data json.load(f) mask np.zeros(shape, dtypenp.uint8) for shape_item in data[shapes]: label shape_item[label] points np.array(shape_item[points], dtypenp.int32) if label tongue_body: class_id 1 # 舌质 elif label tongue_coating: class_id 2 # 舌苔 else: continue cv2.fillPoly(mask, [points], class_id) return mask def process_folder(json_dir, output_dir): os.makedirs(output_dir, exist_okTrue) for file in os.listdir(json_dir): if not file.endswith(.json): continue mask labelme_json_to_mask(os.path.join(json_dir, file)) out_path os.path.join(output_dir, file.replace(.json, .png)) cv2.imwrite(out_path, mask)逻辑说明cv2.fillPoly把多边形顶点填充成指定类别后画的标注如果覆盖前一个会直接覆盖像素值所以在标注时舌质要最后框。这里的class_id顺序要和训练脚本里的ignore_index设置保持一致。处理完的 PNG 是单通道图像像素值 0、1、2。注意不要保存成三通道彩色 PNG否则加载时必须多做一步cv2.COLOR_BGR2GRAY转换且压缩噪声会污染类别索引。3.3 针对舌头图像的数据增强mIoU 能差 6 个点舌头图像有高度统一的成像规范白平衡、光照角度、舌头伸出的程度在不同医疗点差异很大。增强策略需要兼顾几何形变和颜色扰动关键是不要破坏语义边界——舌体是非刚性形变但扭曲太狠会让舌苔纹理失真。import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform A.Compose([ A.RandomResizedCrop(size(512, 512), scale(0.8, 1.0)), A.Rotate(limit15, border_modecv2.BORDER_CONSTANT), A.HorizontalFlip(p0.5), A.OneOf([ A.ColorJitter(brightness0.2, contrast0.2, saturation0.1, hue0.02, p1.0), A.HueSaturationValue(hue_shift_limit5, val_shift_limit20, p1.0), ], p0.8), A.RandomGamma(gamma_limit(80, 120), p0.3), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2(), ]) val_transform A.Compose([ A.Resize(512, 512), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2(), ])参数说明Rotate的border_mode必须设为BORDER_CONSTANT否则旋转后边缘出现的插值像素值会污染背景类在医学分割里这是老坑。RandomResizedCrop的scale下限 0.8 足够太激进会让舌体占不满整张图模型会在背景上学习到不必要的响应。颜色增强里hue的扰动范围控制在 0.02舌色在中医诊断里有临床意义色相漂移过大会让网络把“淡红舌”和“红绛舌”学成同一类。3.4 用 PyTorch Dataset 把图像和掩码配对class TongueDataset(torch.utils.data.Dataset): def __init__(self, image_dir, mask_dir, transformNone): self.image_paths sorted(os.listdir(image_dir)) self.mask_dir mask_dir self.image_dir image_dir self.transform transform def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img_name self.image_paths[idx] img cv2.imread(os.path.join(self.image_dir, img_name)) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask_path os.path.join(self.mask_dir, img_name.replace(.jpg, .png)) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) if self.transform: aug self.transform(imageimg, maskmask) img aug[image] mask aug[mask] mask mask.long() return img, mask逻辑说明文件名用sorted()保证图像和掩码的顺序一致这是 Dataset 实现里最常见的潜藏 bug一旦目录里混入系统隐藏文件或同名不同扩展名的文件顺序全部错位训练指标看起来正常但模型学到的是噪声。如果发现训练集 loss 下降正常、验证集 mIoU 始终不涨先检查__getitem__里返回的图像和掩码是不是同一张。4. 改进版 UNet 训练全流程损失函数、参数配置与模型结构4.1 改进后的 UNet 整体结构代码改进版 UNet 的编码器分四层前两层使用 Res2Block第三层普通卷积第四层瓶颈层再次使用 Res2Block。解码器保持标准结构跳跃连接不额外加注意力模块目的是让对比实验能明确归因于 Res2Net 的贡献。import torch.nn as nn class DownBlock(nn.Module): def __init__(self, in_ch, out_ch, use_res2False, scale4): super().__init__() self.use_res2 use_res2 if use_res2: self.block Res2Block(in_ch, out_ch, scalescale) else: self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) self.pool nn.MaxPool2d(2) def forward(self, x): x self.block(x) return x, self.pool(x) class UNetRes2Net(nn.Module): def __init__(self, in_channels3, num_classes3): super().__init__() self.down1 DownBlock(in_channels, 64, use_res2True) self.down2 DownBlock(64, 128, use_res2True) self.down3 DownBlock(128, 256, use_res2False) self.down4 DownBlock(256, 512, use_res2True, scale4) self.up1 nn.ConvTranspose2d(512, 256, kernel_size2, stride2) self.conv1 nn.Sequential( nn.Conv2d(512, 256, kernel_size3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.Conv2d(256, 256, kernel_size3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue) ) # 后续上采样层级类似省略拼接部分结构说明DownBlock返回两个值一个是当前层输出进跳跃连接一个是池化后的下一层输入。use_res2标志控制哪层替换为 Res2Block。瓶颈层用scale4分支通道是 512/4128这个宽度足够让每个分支学到有区分度的特征。如果scale8每分支只有 64 通道特征碎片化严重收敛变慢实测 mIoU 反而下降 1.2% 左右。4.2 损失函数组合Dice Loss 为主、Focal 补充舌头数据集里背景占比通常超过 60%舌苔和舌质占比加起来约 30%~40%。直接用交叉熵会让背景类主导梯度目标类别的边界预测会非常模糊。实践中效果最稳的组合是Dice Loss Focal Loss权重比 7:3。class DiceLoss(nn.Module): def __init__(self, smooth1.0): super().__init__() self.smooth smooth def forward(self, pred, target): pred torch.softmax(pred, dim1) target_onehot torch.nn.functional.one_hot( target, num_classespred.shape[1] ).permute(0, 3, 1, 2).float() intersection (pred * target_onehot).sum(dim(2, 3)) dice (2.0 * intersection self.smooth) / ( pred.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) self.smooth ) return 1.0 - dice.mean()参数说明smooth是平滑项默认 1.0 是为了防止小目标区域如个别样本里舌苔面积只有几十个像素出现分母为零。one_hot转类别维度的顺序是(B, H, W)到(B, C, H, W)这里permute的维度顺序初学者经常搞混写错后会报维度不匹配训练直接从第一步崩掉。Dice Loss 对前景占比不敏感但对类间边界模糊容忍度低所以配 Focal 来拉低易分样本的权重强化难分的舌苔边界。4.3 一套在 2080Ti 上能跑的训练超参配置用 AdamW 优化器初始学习率 3e-4配合余弦退火。batch size 设为 8512x512 输入如果显存不够优先降低输入尺寸到 448 而不是降低 batch 到 4梯度噪声会明显增大。python train.py \ --arch unet_res2net \ --dataset ./tongue_data \ --image_size 512 \ --batch_size 8 \ --lr 3e-4 \ --loss dicefocal \ --epochs 150 \ --scale 4 \ --seed 42参数说明seed固定为 42 是为了保证对比实验可复现尤其是在验证 Res2Net 改进收益时如果不固定种子两次训练之间 1% 以内的 mIoU 波动会掩盖真实改进。epochs150对小型数据集500~1000 张足够舌头分割不是大模型任务超过 200 epoch 后验证集 mIoU 基本进入平台期继续训练只会增加过拟合风险。评估指标上除了 mIoU 还要专门看Dice coefficient of tongue_coating class。舌苔类别面积小全局 mIoU 可能看起来不错但舌苔类别单独掉到 0.5 以下交给临床用就是废的。训练日志里每 5 个 epoch 打印分类别 IoU这是判断模型是否真的学到了细粒度结构的关键。5. 项目说明书编写要点与模型导出验证一个“可交付”的分割项目代码只占一半分量另一半是项目说明书里的复现信息。下面这套结构是我在多次交付中沉淀下的模板直接按目录写即可。5.1 项目说明书的标准目录结构数据集说明采集设备、标注标准、类别定义、数据划分比例训练/验证/测试 8:1:1环境依赖Python 3.9、PyTorch 1.12、Albumentations 1.3完整 requirements.txt训练步骤数据预处理命令、训练命令、日志输出位置评估结果分模型对比表包含基线 UNet、UNetRes2Net、UNetRes2Net不同 loss 的 mIoU、Dice、参数量复现验证用checkpoint.pth跑推理的命令以及输出结果保存路径复现性最重要的是把随机种子和数据处理版本写清楚。我有一次交付后对方反馈“mIoU 从 84 掉到了 81”最后排查发现是对方用的 OpenCV 版本不同RandomResizedCrop的插值方式变了导致增强分布不一致。项目说明书里务必写上“推荐使用 Docker 镜像”或直接锁死依赖版本大版本号。5.2 用 ONNX 导出并验证推理完整流程导出 ONNX 是部署验证的第一步同时也可以用来检查模型是否在训练和推理模式下行为一致。import torch from models.unet_res2net import UNetRes2Net model UNetRes2Net(in_channels3, num_classes3) model.load_state_dict(torch.load(checkpoints/best.pth, map_locationcpu)) model.eval() dummy torch.randn(1, 3, 512, 512) torch.onnx.export( model, dummy, unet_res2net_tongue.onnx, opset_version12, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}} )逻辑说明dynamic_axes设置动态 batch 维度部署时可以一次推理多张图。opset_version12兼容性好对 BatchNorm 和 Res2Block 里的 split 操作支持稳定。导出后一定要用onnxruntime跑一次推理、对比 PyTorch 输出。5.3 部署时最容易翻车的两个细节第一测试集的图像尺寸必须保持 512 的整数倍或者至少是 16 的倍数。Res2Block 里的torch.chunk按通道切不涉及空间维度但 UNet 下采样四次输入宽高不是 16 的倍数时上采样拼接时特征图尺寸对不上会直接报错。最好在 Dataset 的__getitem__里强制Resize((512, 512))不管原始图多大。第二掩码输出从 logits 转类别时要用torch.argmax在通道维dim1取索引而不是在 softmax 之后取最大概率再转。两者数学上等价但后者的 softmax 计算是浪费的而且fp16推理下 softmax 的精度损失会导致个别像素类别错位。推荐直接对 logits 用argmax。最后一个技巧用poi式的重叠滑窗推理patch-based inference处理超大尺寸舌头图像时重叠率设成 25% 是性能和精度平衡点低于 10% 时拼接缝明显高于 50% 时推理时间翻倍而 mIoU 提升不足 1%。在舌头这类小器官上除非原始图像超过 2048 像素否则直接全图缩放推理即可不需要滑窗。本文还有配套的精品资源点击获取
返回列表