基于 YOLOv8n 与 AttUNet 的胃息肉检测与分割
基于 YOLOv8n 与 AttUnet 的胃息肉检测与分割
代码详见:https://github.com/xiaozhou-alt/Kvasir-SEG
一、项目介绍
胃息肉作为消化道常见病变,其早期精准检测与轮廓分割对临床诊断和治疗方案制定至关重要。传统人工诊断依赖医师经验,存在小息肉漏检、边界定位误差等问题。本项目基于深度学习技术,构建 “检测 - 分割” 二阶段智能分析系统,实现胃息肉的自动定位与精细轮廓提取,为内镜诊断提供辅助工具。
二、文件夹结构
Kvasir-SEG/
├── data/ # 数据目录
├── Kvasir-SEG/ # 原始数据集(需自行下载放置)
├── images/
├── bbox/
└── masks/
└── dataset/ # 检测任务处理后数据(代码自动生成)
├── images/
├── train/
└── val/
├── labels/
├── train/
└── val/
└── data.yaml
├── train_pre.py # 息肉检测代码(YOLOv8)
├── train_mask.py # 息肉分割代码(AttUNet)
├── log/
├── output/ # 输出结果目录(代码自动生成)
├── polyp_detection/ # 检测模型输出
└── exp/
└── weights/ # 最佳检测模型权重(best.pt)
├── model/
└── pic/
├── requirements.txt
└── README.md
三、数据集介绍
1 数据集来源
采用公开医学影像数据集 Kvasir-SEG(可通过 Kaggle 获取:Kvasir-SEG Data (Polyp segmentation & detection) (kaggle.com)),该数据集包含临床内镜下胃息肉图像及对应标注,是息肉分割任务的基准数据集。
2 数据结构
原始数据集包含两个核心目录,适配检测与分割双任务:
Kvasir-SEG/
├── images/ # 原始内镜图像(支持jpg格式)
├── bbox/ # 检测任务标注(CSV格式,含类别、边界框坐标)
└── masks/ # 分割任务标注(灰度图像,息肉区域为白色)
3 数据处理细节
-
检测任务:
- 校验图像有效性(排除空图、损坏文件);
- 将 CSV 标注(class_name, xmin, ymin, xmax, ymax) 转换为 YOLO 格式(归一化中心坐标 + 宽高);
- 按 8:2 比例随机划分训练集 / 验证集,生成
data.yaml配置文件。
.
-
分割任务:
- 图像与掩码尺寸统一 resize 至 256 × 256 256×256 256×256;
- 按 85 : 15 85:15 85:15 比例划分训练集 / 验证集,小样本时自动重复数据增强训练;
- 内置多维度数据增强(翻转、旋转、缩放、弹性变换等)。
四、YOLOv8n 与 AttUnet 模型介绍
1. YOLOv8模型:高效息肉检测的核心
YOLOv8 采用 无锚点(Anchor-Free) 设计,通过端到端的方式实现目标检测,在速度和精度上均有显著提升。结合本项目息肉检测的需求,其核心模块可拆解为:输入端预处理、主干网络特征提取、检测头与损失函数设计。

1.1 输入端预处理模块
YOLOv8 的输入端预处理包含图像尺寸调整、归一化、Mosaic 数据增强等操作。对于输入图像 I ∈ R H × W × 3 I \in \mathbb{R}^{H \times W \times 3} I∈RH×W×3,首先将其 Resize 至固定尺寸 I r e s i z e ∈ R 640 × 640 × 3 I_{resize} \in \mathbb{R}^{640 \times 640 \times 3} Iresize∈R640×640×3(本项目配置 imgsz=640),随后进行归一化处理,将像素值映射至 [ 0 , 1 ] [0,1] [0,1] 区间,公式如下:
I n o r m ( i , j , c ) = I r e s i z e ( i , j , c ) 255.0 I_{norm}(i,j,c) = \frac{I_{resize}(i,j,c)}{255.0} Inorm(i,j,c)=255.0Iresize(i,j,c)
同时,为提升模型泛化能力,采用 Mosaic 增强(项目中 mosaic=1.0),通过随机裁剪 4 4 4 张图像并拼接生成新样本,增强模型对不同尺度、位置息肉的适应能力。此外,项目代码中还实现了随机翻转(fliplr=0.5)、平移(translate=0.1)等增强策略。
🤓🤓🤓小周有话说:
这一步就像我们在给医生准备息肉检查的图像资料。
首先把不同大小的内镜图像统一调整成 640 × 640 640×640 640×640 的标准尺寸,方便后续统一分析;然后把图像的像素亮度值从 0 − 255 0-255 0−255 的范围缩小到 0 − 1 0-1 0−1,就像把测量单位从“厘米”换成“米”,让模型计算更高效。
Mosaic 增强则像是 把 4 张不同的内镜图随机剪几块拼在一起,比如把有小息肉、大息肉、息肉在边缘、息肉在中心的图像拼合成新图,这样模型就能见多识广,遇到各种位置和大小的息肉都能准确识别。
1.2 主干网络:特征提取核心
YOLOv8的主干网络采用 C2f 模块 和 SPPF(Spatial Pyramid Pooling - Fast)模块 构建。C2f 模块 基于 残差连接 思想,通过分流和融合实现特征的高效提取,其核心是将输入特征图 F i n ∈ R C × H × W F_{in} \in \mathbb{R}^{C \times H \times W} Fin∈RC×H×W 分为两部分,一部分直接通过,另一部分经过多组瓶颈层(Bottleneck)处理后与前一部分融合,输出特征图 F o u t ∈ R 2 C × H × W F_{out} \in \mathbb{R}^{2C \times H \times W} Fout∈R2C×H×W(默认配置下),融合公式如下:
F o u t = C o n c a t ( F i n 1 , B o t t l e n e c k ( F i n 2 ) ) F_{out} = Concat(F_{in1}, Bottleneck(F_{in2})) Fout=Concat(Fin1,Bottleneck(Fin2))
其中 F i n 1 F_{in1} Fin1 和 F i n 2 F_{in2} Fin2 是 F i n F_{in} Fin 的分流结果。SPPF 模块 则通过不同尺寸的池化核 对特征图进行池化,再将结果拼接,增强模型对多尺度特征的捕捉能力,其输出特征图维度与输入保持一致。
🤓🤓🤓小周有话说:
主干网络就像一个 “特征筛选器”,专门从内镜图像中提取和息肉相关的关键信息。
C2f 模块 的分流融合机制,就像我们筛选有用信息时,一部分信息直接保留,另一部分信息经过深度分析后再整合,这样既能保证信息的完整性,又能提升有用信息的浓度。比如图像中的息肉边缘、纹理、颜色等特征,经过 C2f 模块处理后会被强化。
而 SPPF 模块 则像是用不同倍数的放大镜观察图像,小放大镜找小息肉的细节,大放大镜抓大息肉的整体轮廓,最后把不同放大镜看到的信息结合起来,让模型能精准捕捉不同大小的息肉特征。
1.3 检测头与损失函数
YOLOv8 采用 Anchor-Free检测头,直接预测目标的中心坐标、宽高和置信度。对于主干网络输出的 3 3 3 个不同尺度特征图 F 1 ∈ R C 1 × H 1 × W 1 F_1 \in \mathbb{R}^{C1 \times H1 \times W1} F1∈RC1×H1×W1、 F 2 ∈ R C 2 × H 2 × W 2 F_2 \in \mathbb{R}^{C2 \times H2 \times W2} F2∈RC2×H2×W2、 F 3 ∈ R C 3 × H 3 × W 3 F_3 \in \mathbb{R}^{C3 \times H3 \times W3} F3∈RC3×H3×W3,通过检测头映射为预测张量 P 1 ∈ R ( 4 + 1 + n c ) × H 1 × W 1 P_1 \in \mathbb{R}^{(4+1+nc) \times H1 \times W1} P1∈R(4+1+nc)×H1×W1、 P 2 ∈ R ( 4 + 1 + n c ) × H 2 × W 2 P_2 \in \mathbb{R}^{(4+1+nc) \times H2 \times W2} P2∈R(4+1+nc)×H2×W2、 P 3 ∈ R ( 4 + 1 + n c ) × H 3 × W 3 P_3 \in \mathbb{R}^{(4+1+nc) \times H3 \times W3} P3∈R(4+1+nc)×H3×W3,其中 4 4 4 代表目标坐标( x c , y c , w , h x_c, y_c, w, h xc,yc,w,h), 1 1 1 代表置信度, n c nc nc 为类别数(本项目nc=1,仅息肉类)。
损失函数 采用 CIoU损失(坐标损失)+ BCEWithLogitsLoss(置信度损失)+ CrossEntropyLoss(类别损失)的组合形式,坐标损失公式如下:
L C I o U = 1 − I o U + ρ 2 ( b , b g t ) c 2 + α v L_{CIoU} = 1 - IoU + \frac{\rho^2(b, b^{gt})}{c^2} + \alpha v LCIoU=1−IoU+c2ρ2(b,bgt)+αv
其中 b b b 和 b g t b^{gt} bgt 分别为 预测框 和 真实框坐标, ρ \rho ρ 为欧氏距离, c c c 为包围两框的最小矩形对角线长度, α \alpha α 为权重系数, v v v 为衡量框的长宽比一致性的参数。
🤓🤓🤓小周有话说:
检测头就像模型的 “识别和定位工具”,专门负责在处理后的特征图上找息肉、标位置。
Anchor-Free 设计就像医生直接用手在图像上圈出息肉,而不是先准备好固定大小的框去套。比如看到一个息肉,检测头会直接预测它的中心位置(像地图上的经纬度)、宽度和高度,同时给出“这是息肉”的置信度。
损失函数 则是模型的 “纠错老师”,如果模型把息肉的位置圈偏了、大小估错了,或者把正常组织当成了息肉,损失函数就会给出“惩罚”,让模型下次调整得更准确。比如模型预测的息肉框和真实息肉框偏差越大,CIoU 损失就越大,模型就知道需要大幅调整框的位置。
2. AttUnet模型:精准息肉分割的关键
AttUnet 在经典 Unet 的基础上引入 注意力门控(Attention Gate)模块,解决了 Unet 在分割时对无关区域特征过度关注的问题,大幅提升了息肉分割的精准度。其核心模块包括:编码-解码结构、注意力门控模块、损失函数设计。
如下为 AttUNet 网络架构图:

2.1 编码-解码核心结构
AttUnet 采用 对称的编码-解码结构。编码端(下采样)由 DoubleConv模块 和 MaxPool2d 组成,DoubleConv 模块包含两组 Conv2d + BatchNorm2d + ReLU + Dropout 的组合,对输入特征图进行深度特征提取,公式如下:
F c o n v = R e L U ( B a t c h N o r m 2 d ( C o n v 2 d ( F i n , o u t _ c h a n n e l s , k e r n e l _ s i z e = 3 , p a d d i n g = 1 ) ) ) F_{conv} = ReLU(BatchNorm2d(Conv2d(F_{in}, out\_channels, kernel\_size=3, padding=1))) Fconv=ReLU(BatchNorm2d(Conv2d(Fin,out_channels,kernel_size=3,padding=1)))
通过 MaxPool2d(步长2)实现下采样,将特征图尺寸减半、通道数加倍,逐步提升特征的语义信息。解码端(上采样)采用Upsample(双线性插值)实现特征图尺寸加倍,再与编码端对应尺度的特征图(经注意力门控处理)拼接(Concat),随后通过 DoubleConv 模块融合特征。最终通过 1 × 1 1×1 1×1 卷积将通道数映射为类别数(本项目 n_classes=1),得到分割结果。
特征图尺寸变化规律:设输入图像尺寸为 H × W × 3 H \times W \times 3 H×W×3(本项目 Resize 至 256 × 256 256×256 256×256),编码端经过 4 4 4次下采样后,特征图尺寸变为 H / 16 × W / 16 H/16 \times W/16 H/16×W/16;解码端经过 4 4 4 次上采样后,最终输出尺寸恢复为 H × W H \times W H×W。
🤓🤓🤓小周有话说:
AttUnet 的 编码-解码结构 就像 “图像拼图+细节补充”的过程。
编码端 就像我们把一张完整的内镜图像逐步缩小,每缩小一次就更关注图像的整体特征(比如息肉的大致区域),同时忽略一些无关的细节(比如图像上的微小噪声)。比如 256 × 256 256×256 256×256 的图像经过 4 4 4 次下采样后,变成 16 × 16 16×16 16×16 的特征图,此时里面包含的是息肉的核心语义信息。
解码端 则是把缩小的特征图逐步放大,每次放大时,都会从编码端拿对应的细节特征(经过注意力门控筛选后)补充进来,就像拼图时先拼出大致轮廓,再逐步填充细节。最终恢复到 256 × 256 256×256 256×256 的尺寸,精准勾勒出息肉的边界。
2.2 注意力门控(Attention Gate)模块
注意力门控模块 是 AttUnet 的核心创新点,其作用是抑制背景区域特征,强化目标区域(息肉)特征。模块输入包括编码端特征 x ∈ R C x × H × W x \in \mathbb{R}^{C_x \times H \times W} x∈RCx×H×W(输入特征)和解码端特征 g ∈ R C g × H g × W g g \in \mathbb{R}^{C_g \times H_g \times W_g} g∈RCg×Hg×Wg(门控信号),首先通过 1 × 1 1×1 1×1 卷积将两者映射到相同通道数 C i n t e r C_{inter} Cinter:
x t r a n s = C o n v 2 d ( x , C i n t e r , k e r n e l _ s i z e = 1 ) , g g a t e = C o n v 2 d ( g , C i n t e r , k e r n e l _ s i z e = 1 ) x_{trans} = Conv2d(x, C_{inter}, kernel\_size=1), \quad g_{gate} = Conv2d(g, C_{inter}, kernel\_size=1) xtrans=Conv2d(x,Cinter,kernel_size=1),ggate=Conv2d(g,Cinter,kernel_size=1)
若尺寸不匹配,通过上采样或下采样调整至一致,随后进行元素相加并经过 R e L U ReLU ReLU 激活,再通过 1 × 1 1×1 1×1 卷积和 Sigmoid 激活得到注意力权重 ψ ∈ R 1 × H × W \psi \in \mathbb{R}^{1 \times H \times W} ψ∈R1×H×W:
ψ = S i g m o i d ( C o n v 2 d ( R e L U ( x t r a n s + g g a t e ) , 1 , k e r n e l _ s i z e = 1 ) ) \psi = Sigmoid(Conv2d(ReLU(x_{trans} + g_{gate}), 1, kernel\_size=1)) ψ=Sigmoid(Conv2d(ReLU(xtrans+ggate),1,kernel_size=1))
最终将注意力权重与编码端特征相乘,得到增强后的目标特征: x a t t = x × ψ x_{att} = x \times \psi xatt=x×ψ。
🤓🤓🤓小周有话说:
注意力门控模块 就像模型的 “聚焦镜”,专门让模型把注意力集中在息肉区域,忽略正常的胃肠壁组织。比如在处理内镜图像时,编码端会提取到很多特征,其中既有息肉的特征,也有正常黏膜的特征。
当这些特征传递到解码端时,注意力门控模块会根据解码端的门控信号(大致的息肉区域信息),给不同的特征打分:息肉相关的特征打高分(权重接近1),正常组织的特征打低分(权重接近0)。然后用这个分数去筛选编码端的特征,只把高分的息肉特征传递给解码端。就像医生看内镜图时,会自动聚焦在息肉上,忽略周围的正常组织,从而更精准地判断息肉的范围。
2.3 损失函数设计
考虑到息肉分割任务中,息肉区域可能占比较小(类别不平衡),本项目采用组合损失函数:BCEWithLogitsLoss(交叉熵损失)+ DiceLoss,公式如下:
L c o m b i n e d = ω b c e × L b c e + ω d i c e × L d i c e L_{combined} = \omega_{bce} \times L_{bce} + \omega_{dice} \times L_{dice} Lcombined=ωbce×Lbce+ωdice×Ldice
其中 ω b c e = 0.3 \omega_{bce}=0.3 ωbce=0.3、 ω d i c e = 0.7 \omega_{dice}=0.7 ωdice=0.7 为权重系数。交叉熵损失 L b c e L_{bce} Lbce 衡量预测概率与真实标签的差异,DiceLoss基于Dice系数计算,公式为:
L d i c e = 1 − 2 × ∣ X ∩ Y ∣ ∣ X ∣ + ∣ Y ∣ L_{dice} = 1 - \frac{2 \times |X \cap Y|}{|X| + |Y|} Ldice=1−∣X∣+∣Y∣2×∣X∩Y∣
其中 X X X为预测分割结果, Y Y Y为真实分割标签, ∣ ⋅ ∣ |·| ∣⋅∣表示像素数量。DiceLoss 对小目标区域更敏感,能有效提升息肉区域的分割精度。
🤓🤓🤓小周有话说:
AttUnet 的 损失函数 就像“双重纠错老师”,专门解决息肉区域小、容易分割不准的问题。
交叉熵损失 负责判断每个像素是不是息肉,比如把息肉像素判成正常像素,就会给出惩罚;
D i c e L o s s DiceLoss DiceLoss 则负责衡量模型分割的息肉区域和真实息肉区域的重叠度,重叠度越低,惩罚越重。
因为息肉在图像中可能只占很小一块,单一的交叉熵损失可能会让模型“偏向”预测大部分区域为正常组织,而 DiceLoss 能强制模型关注小的息肉区域。
两者结合,就能让模型既准确判断每个像素的类别,又能精准勾勒出息肉的完整轮廓。
五、项目实现
1. 评估指标与损失函数模块
1.1 分割评估指标计算
# 评估指标计算函数(保持不变)
def calculate_metrics(pred, target, threshold=0.5):
"""
计算分割任务的评估指标
pred: 模型预测输出 (经过sigmoid)
target: 真实标签
threshold: 二值化阈值
"""
# 将预测值二值化
pred = (pred > threshold).float()
# 计算TP, TN, FP, FN
TP = (pred * target).sum()
TN = ((1 - pred) * (1 - target)).sum()
FP = (pred * (1 - target)).sum()
FN = ((1 - pred) * target).sum()
# 计算精确率
precision = TP / (TP + FP + 1e-8) # 加小值避免除零
# 计算召回率
recall = TP / (TP + FN + 1e-8)
# 计算Dice系数
dice = (2 * TP) / (2 * TP + FP + FN + 1e-8)
# 计算IoU (交并比)
iou = TP / (TP + FP + FN + 1e-8)
return {
'precision': precision.item(),
'recall': recall.item(),
'dice': dice.item(),
'iou': iou.item()
}
- 适用场景:用于评估胃息肉分割模型的预测效果,核心指标覆盖分割任务的关键评价维度。
- 输入要求:
pred为模型输出经过 sigmoid 激活后的概率图(值在 0 − 1 0-1 0−1 之间),target为真实掩码( 0 / 1 0/1 0/1 二值图)。 - 核心逻辑:
- 二值化:将预测概率图按阈值(默认 0.5 0.5 0.5)转换为二值图( 0 0 0 表示背景, 1 1 1 表示息肉)。
- 混淆矩阵元素计算:通过张量点乘计算真阳性( T P TP TP,预测息肉且真实为息肉)、真阴性( T N TN TN,预测背景且真实为背景)、假阳性( F P FP FP,预测息肉但真实为背景)、假阴性( F N FN FN,预测背景但真实为息肉)。
- 指标计算:
- 精确率( P r e c i s i o n Precision Precision):预测为息肉的样本中真实为息肉的比例,衡量预测准确性。
- 召回率( R e c a l l Recall Recall):真实为息肉的样本中被正确预测的比例,衡量息肉区域的覆盖能力。
- D i c e Dice Dice 系数:衡量预测区域与真实区域的重叠度,取值 0 − 1 0-1 0−1,越接近 1 1 1 表示重叠度越高。
- I o U IoU IoU(交并比):预测区域与真实区域的交集除以并集,直观反映分割区域的匹配程度。
- 数值稳定性:所有除法都添加 1 e − 8 1e-8 1e−8,避免分母为 0 0 0 导致的计算错误。
- 输出:返回包含 4 4 4 个指标的字典,且通过
.item()将张量转换为 P y t h o n Python Python 数值,便于后续记录和可视化。
1.2 Dice损失函数
# 自定义Dice损失函数(保持不变)
class DiceLoss(nn.Module):
def __init__(self, smooth=1e-5):
super(DiceLoss, self).__init__()
self.smooth = smooth
def forward(self, input, target):
input = torch.sigmoid(input)
intersection = (input * target).sum()
return 1 - (2. * intersection + self.smooth) / (input.sum() + target.sum() + self.smooth)
- 设计思路:基于 Dice系数 的损失函数,
Dice系数越接近 1 表示分割效果越好,因此损失函数定义为 1 − D i c e 1 - Dice 1−Dice 系数,使损失越小分割效果越好。 - 关键步骤:
- 激活函数:输入
input为模型未经过 sigmoid 的原始输出,需先通过 sigmoid 转换为 0 − 1 0-1 0−1 的概率值。 - 交集计算:通过张量点乘
input * target得到预测与真实掩码的交集,求和得到交集大小。 - 损失计算:引入
smooth(默认 1 e − 5 1e-5 1e−5)避免分子/分母为 0 0 0,提升数值稳定性。
- 激活函数:输入
- 相比交叉熵损失,Dice 损失更关注分割区域的重叠度,对类别不平衡(如息肉区域占比小)场景更友好,能有效避免模型偏向预测占比大的背景类。
1.3 混合损失函数
# 混合损失函数(交叉熵+Dice)(保持不变)
class CombinedLoss(nn.Module):
def __init__(self, bce_weight=0.5, dice_weight=0.5, smooth=1e-5):
super(CombinedLoss, self).__init__()
self.bce_loss = nn.BCEWithLogitsLoss()
self.dice_loss = DiceLoss(smooth)
self.bce_weight = bce_weight
self.dice_weight = dice_weight
def forward(self, input, target):
bce = self.bce_loss(input, target)
dice = self.dice_loss(input, target)
return self.bce_weight * bce + self.dice_weight * dice
- 设计思路:融合 二元交叉熵损失(BCEWithLogitsLoss) 和 Dice损失 的优势,兼顾 类别概率分布 的拟合(交叉熵)和 分割区域的重叠度(Dice),提升模型的分割性能。
- 组件说明:
bce_loss:使用BCEWithLogitsLoss,直接接收模型原始输出(无需提前 sigmoid),计算预测概率与真实标签的交叉熵,擅长拟合类别分布。dice_loss:调用自定义的DiceLoss,关注区域重叠度。
- 权重调节:通过
bce_weight和dice_weight(默认各 0.5 0.5 0.5)调节两种损失的贡献比例,可根据数据集特点(如息肉大小、不平衡程度)调整,后续主函数中设置为 B C E BCE BCE 权重 0.3 0.3 0.3、 D i c e Dice Dice 权重 0.7$,更侧重区域重叠度。 - 解决单一损失的局限性,交叉熵保证模型的梯度稳定性, D i c e Dice Dice 损失解决类别不平衡问题,两者结合使模型在训练过程中更稳定,分割精度更高。
2. 数据增强部分
2.1 随机高斯模糊
# 增强的数据增强策略
class RandomGaussianBlur(object):
def __init__(self, p=0.5, radius_range=(0.5, 2.0)):
self.p = p
self.radius_range = radius_range
def __call__(self, img):
if random.random() < self.p:
radius = random.uniform(*self.radius_range)
return img.filter(ImageFilter.GaussianBlur(radius=radius))
return img
- 功能:模拟图像拍摄时的模糊场景,增加数据多样性,提升模型对模糊图像的鲁棒性。
- 参数说明:
p:执行该增强的概率(默认 0.5 0.5 0.5),即 50 % 50\% 50% 的图像会被模糊处理。radius_range:高斯模糊核半径的范围(默认 0.5 − 2.0 0.5-2.0 0.5−2.0),半径越大模糊程度越高。
- 执行逻辑:调用时生成随机数,若小于
p则随机选择半径进行高斯模糊,否则返回原始图像
2.2 随机弹性变换(修复维度不匹配)
# 修复的弹性变换类 - 解决维度不匹配问题
class RandomElasticTransform(object):
def __init__(self, p=0.4, alpha=120, sigma=15):
self.p = p
self.alpha = alpha
self.sigma = sigma
def __call__(self, img):
if random.random() < self.p:
img = np.array(img)
# 获取图像尺寸(高度、宽度),忽略通道数
h, w = img.shape[:2]
# 创建2D位移场(与图像尺寸匹配)
dx = cv2.GaussianBlur((np.random.rand(h, w) * 2 - 1),
(0, 0), self.sigma) * self.alpha
dy = cv2.GaussianBlur((np.random.rand(h, w) * 2 - 1),
(0, 0), self.sigma) * self.alpha
# 创建坐标网格
x, y = np.meshgrid(np.arange(w), np.arange(h))
map_x = (x + dx).astype(np.float32)
map_y = (y + dy).astype(np.float32)
# 对每个通道应用相同的变换
if len(img.shape) == 3: # 彩色图像
transformed = np.zeros_like(img)
for c in range(img.shape[2]):
transformed[:, :, c] = cv2.remap(
img[:, :, c], map_x, map_y,
interpolation=cv2.INTER_LINEAR,
borderMode=cv2.BORDER_REFLECT
)
return Image.fromarray(transformed)
else: # 灰度图像
img = cv2.remap(
img, map_x, map_y,
interpolation=cv2.INTER_LINEAR,
borderMode=cv2.BORDER_REFLECT
)
return Image.fromarray(img)
return img
- 功能:模拟组织形变(如胃壁蠕动导致的息肉形态变化),增强模型对息肉不同形态的适应能力,核心修复了原始弹性变换中可能出现的维度不匹配问题。
- 参数说明:
p:执行概率(默认 0.4 0.4 0.4)。alpha:控制位移幅度(默认 120 120 120),值越大形变越剧烈。sigma:控制位移场的平滑度(默认 15 15 15),值越大形变越连贯。
- 核心逻辑(修复维度不匹配的关键):
- 先将
PIL图像转换为numpy数组,提取图像的 高度( h h h) 和 宽度( w w w),忽略通道数,确保位移场尺寸与图像尺寸严格匹配。 - 生成
2D位移场(dx、dy):通过随机生成 [ − 1 , 1 ] [-1,1] [−1,1] 的矩阵,经高斯模糊平滑后乘以 alpha 得到最终位移,保证位移场是与图像尺寸一致的 2D 矩阵。 - 坐标网格与重映射:创建像素坐标网格 ( x 、 y ) (x、y) (x、y),叠加位移场得到新坐标 ( m a p _ x 、 m a p _ y ) (map\_x、map\_y) (map_x、map_y);通过
cv2.remap对每个通道(彩色图像)或单通道(灰度图像)进行 重映射,实现弹性形变。
- 先将
- 边界处理:使用
BORDER_REFLECT(反射边界),避免形变后图像边缘出现黑边,保证图像完整性。
2.3 随机缩放
class RandomZoom(object):
def __init__(self, p=0.5, zoom_range=(0.8, 1.2)):
self.p = p
self.zoom_range = zoom_range
def __call__(self, img):
if random.random() < self.p:
zoom = random.uniform(*self.zoom_range)
w, h = img.size
new_w, new_h = int(w * zoom), int(h * zoom)
img = img.resize((new_w, new_h), Image.BILINEAR)
# 如果缩小了,随机裁剪回原尺寸
if zoom < 1.0:
x = random.randint(0, new_w - w) if new_w > w else 0
y = random.randint(0, new_h - h) if new_h > h else 0
img = img.crop((x, y, x + w, y + h))
# 如果放大了,中心裁剪
elif zoom > 1.0:
x = (new_w - w) // 2
y = (new_h - h) // 2
img = img.crop((x, y, x + w, y + h))
return img
- 功能:模拟不同拍摄距离下的息肉大小变化,增强模型对息肉不同尺度的识别能力,确保缩放后图像尺寸与原始一致(便于批量训练)。
- 参数说明:
p:执行概率(默认 0.5 0.5 0.5)。zoom_range:缩放比例范围(默认 0.8 − 1.2 0.8-1.2 0.8−1.2),即图像可缩小至 80 % 80\% 80% 或放大至 120 % 120\% 120%。
- 执行逻辑:
- 随机选择缩放比例,先将图像 resize 到新尺寸(使用 双线性插值
Image.BILINEAR,保证图像质量)。 - 缩放后裁剪回原尺寸:
- 缩小(zoom<1.0):新尺寸大于原尺寸,随机选择裁剪区域(确保裁剪后尺寸与原尺寸一致),增加随机性。
- 放大(zoom>1.0):新尺寸大于原尺寸,中心裁剪(保证息肉主体不被裁掉)。
- 随机选择缩放比例,先将图像 resize 到新尺寸(使用 双线性插值
- 避免缩放后尺寸不一致导致的批量训练错误,同时通过随机/中心裁剪平衡随机性和主体保留需求。
3. 数据集与数据加载模块
# 数据集类定义(保持不变)
class PolypDataset(Dataset):
def __init__(self, image_paths, mask_paths, transform=None, is_train=True):
self.image_paths = image_paths
self.mask_paths = mask_paths
self.transform = transform
self.is_train = is_train
# 数据重复因子,小样本时增加训练次数
self.repeat = 2 if is_train and len(image_paths) < 1500 else 1
def __len__(self):
return len(self.image_paths) * self.repeat
def __getitem__(self, idx):
# 处理重复索引
idx = idx % len(self.image_paths)
# 加载图像和掩码
image = Image.open(self.image_paths[idx]).convert('RGB')
mask = Image.open(self.mask_paths[idx]).convert('L')
# 更多数据增强
if self.is_train:
# 随机对比度增强
if random.random() < 0.3:
factor = random.uniform(0.7, 1.3)
image = ImageEnhance.Contrast(image).enhance(factor)
# 随机亮度增强
if random.random() < 0.3:
factor = random.uniform(0.7, 1.3)
image = ImageEnhance.Brightness(image).enhance(factor)
# 应用变换
if self.transform:
image = self.transform['train'](image)
mask = self.transform['mask'](mask)
# 将掩码二值化
mask = (mask > 0.5).float()
return image, mask
- 类定位:继承 P y T o r c h PyTorch PyTorch 的
Dataset类,自定义适用于胃息肉分割任务的数据集,实现图像与对应掩码的加载、增强和格式转换。 - __init__方法(初始化):
- 输入参数:图像路径列表(
image_paths)、掩码路径列表(mask_paths)、变换字典(transform)、是否为训练集(is_train)。 - 重复因子(
repeat):针对小样本场景(训练集样本数< 1500 1500 1500),将训练集重复 2 2 2 次,增加训练迭代次数,提升模型泛化能力;验证集不重复。
- 输入参数:图像路径列表(
- __len__方法(数据集长度):返回样本数乘以重复因子,决定训练时的迭代次数。
- __getitem__方法(核心,获取单样本):
- 索引处理:通过
idx % len(self.image_paths)解决重复因子导致的索引超出原始样本数的问题,循环使用原始样本。 - 图像与掩码加载:
- 图像:用 P I L PIL PIL 打开,转换为 R G B RGB RGB 模式( 3 3 3 通道),符合模型输入要求。
- 掩码:用 P I L PIL PIL 打开,转换为 L L L 模式(单通道灰度图),对应息肉(非 0 0 0 值)和背景( 0 0 0 值)。
- 训练集专属增强:在应用预设变换前,额外添加随机对比度( 30 % 30\% 30% 概率,对比度因子 0.7 − 1.3 0.7-1.3 0.7−1.3)和亮度增强( 30 % 30\% 30% 概率,亮度因子 0.7 − 1.3 0.7-1.3 0.7−1.3),进一步丰富训练数据多样性。
- 变换应用:通过
transform字典分别对图像(train变换链)和掩码(mask变换链)进行处理(如resize、随机翻转、归一化等)。 - 掩码二值化:将掩码转换为 0 / 1 0/1 0/1 二值张量(阈值 0.5 0.5 0.5),符合分割任务的标签要求。
- 索引处理:通过
- 输出:返回处理后的图像张量和对应的二值掩码张量,用于模型训练/验证。
4. 注意力U-Net模型模块
4.1 注意力门模块
# 注意力模块(保持不变)
class AttentionGate(nn.Module):
def __init__(self, gate_channels, input_channels, inter_channels=None):
super(AttentionGate, self).__init__()
if inter_channels is None:
inter_channels = input_channels // 2
self.gate = nn.Sequential(
nn.Conv2d(gate_channels, inter_channels, kernel_size=1, stride=1, padding=0),
nn.BatchNorm2d(inter_channels)
)
self.transform = nn.Sequential(
nn.Conv2d(input_channels, inter_channels, kernel_size=1, stride=1, padding=0),
nn.BatchNorm2d(inter_channels)
)
self.psi = nn.Sequential(
nn.Conv2d(inter_channels, 1, kernel_size=1, stride=1, padding=0),
nn.BatchNorm2d(1),
nn.Sigmoid()
)
self.relu = nn.ReLU(inplace=True)
# 上采样层,用于匹配尺寸
self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
def forward(self, x, g):
# x: 编码器特征图 (输入特征)
# g: 解码器特征图 (门控信号)
g_conv = self.gate(g)
x_conv = self.transform(x)
# 确保g_conv和x_conv尺寸匹配
if g_conv.size()[2:] != x_conv.size()[2:]:
# 根据需要进行上采样或下采样以匹配尺寸
if g_conv.size()[2] < x_conv.size()[2]:
g_conv = self.upsample(g_conv)
else:
x_conv = nn.functional.adaptive_max_pool2d(x_conv, g_conv.size()[2:])
# 相加并激活
psi = self.relu(g_conv + x_conv)
psi = self.psi(psi)
# 应用注意力权重
return x * psi
- 核心作用:作为 U − N e t U-Net U−Net 解码器与编码器特征融合的桥梁,通过门控信号(解码器特征)引导编码器特征的注意力加权,使模型聚焦于息肉区域的关键特征,抑制背景噪声,提升分割精度。该模块与 Y O L O YOLO YOLO 检测模型形成互补—— Y O L O YOLO YOLO 负责快速定位息肉大致区域(边界框),注意力模块则助力分割模型精准捕捉息肉边缘细节。
- __init__方法(模块初始化):
- 输入参数:门控信号通道数(
gate_channels,解码器特征通道数)、输入特征通道数(input_channels,编码器特征通道数)、中间通道数(inter_channels,默认输入通道数的 1 / 2 1/2 1/2,降低计算量)。 - 三大分支:
- 门控分支(
gate): 1 × 1 1×1 1×1 卷积+批归一化,将解码器特征映射到中间通道数,提取门控信号。 - 变换分支(
transform): 1 × 1 1×1 1×1 卷积+批归一化,将编码器特征映射到中间通道数,与门控信号维度匹配。 - psi分支(
psi): 1 × 1 1×1 1×1 卷积+批归一化 + sigmoid,将融合特征转换为 0 − 1 0-1 0−1 的注意力权重图。
- 门控分支(
- 辅助组件: R e L U ReLU ReLU 激活函数(引入非线性)、上采样层(用于尺寸匹配)。
- 输入参数:门控信号通道数(
- forward 方法(前向传播):
- 输入:
x(编码器特征图)、g(解码器特征图,门控信号)。 - 特征映射:通过
gate和transform分支分别处理门控信号和编码器特征。 - 尺寸匹配(关键改进):若两者尺寸不一致,通过上采样(门控信号尺寸小时)或自适应最大池化(编码器特征尺寸小时)调整,确保后续能直接相加,解决 U − N e t U-Net U−Net 不同层级特征尺寸不匹配的问题。
- 注意力权重生成:将匹配后的特征相加,经 R e L U ReLU ReLU 激活后通过
psi分支生成注意力权重图。 - 权重应用:将编码器特征与注意力权重图逐元素相乘,强化关键区域特征,抑制无关背景特征。此过程可结合 Y O L O YOLO YOLO 检测出的息肉边界框区域,进一步聚焦有效区域,减少无效背景特征的干扰。
- 输入:
- 输出:经过注意力加权后的编码器特征,用于后续与解码器特征拼接融合。
4.2 U-Net基础组件(双重卷积、下采样、上采样、输出卷积)
# U-Net模型组件(保持不变)
class DoubleConv(nn.Module):
def __init__(self, in_channels, out_channels, dropout=0.2):
super(DoubleConv, self).__init__()
self.double_conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True),
nn.Dropout2d(dropout),
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
def forward(self, x):
return self.double_conv(x)
class Down(nn.Module):
def __init__(self, in_channels, out_channels, dropout=0.2):
super(Down, self).__init__()
self.maxpool_conv = nn.Sequential(
nn.MaxPool2d(2),
DoubleConv(in_channels, out_channels, dropout)
)
def forward(self, x):
return self.maxpool_conv(x)
class Up(nn.Module):
def __init__(self, in_channels, out_channels, skip_channels, bilinear=True, dropout=0.2):
super(Up, self).__init__()
self.attention = AttentionGate(
gate_channels=in_channels // 2 if not bilinear else in_channels // 2,
input_channels=skip_channels
)
if bilinear:
self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
self.conv = DoubleConv(in_channels, out_channels, dropout)
else:
self.up = nn.ConvTranspose2d(in_channels//2, in_channels//2, kernel_size=2, stride=2)
self.conv = DoubleConv(in_channels, out_channels, dropout)
def forward(self, x1, x2):
x2 = self.attention(x2, x1)
x1 = self.up(x1)
# 输入尺寸调整
diffY = x2.size()[2] - x1.size()[2]
diffX = x2.size()[3] - x1.size()[3]
x1 = nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2,
diffY // 2, diffY - diffY // 2])
x = torch.cat([x2, x1], dim=1)
return self.conv(x)
class OutConv(nn.Module):
def __init__(self, in_channels, out_channels):
super(OutConv, self).__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)
def forward(self, x):
return self.conv(x)
-
DoubleConv(双重卷积模块):
- 功能: U − N e t U-Net U−Net 的核心特征提取单元,通过两次 3 × 3 3×3 3×3 卷积(
padding=1,保证尺寸不变)提取图像特征。该模块提取的细粒度特征可与 Y O L O YOLO YOLO 检测的粗粒度定位信息结合,提升对微小息肉的识别能力。 - 结构:卷积+批归一化(加速训练、防止过拟合)+ R e L U ReLU ReLU(非线性激活)+ d r o p o u t dropout dropout(随机失活,防止过拟合)+再卷积+批归一化 + R e L U ReLU ReLU。
- 输入输出:输入通道数
in_channels,输出通道数out_channels,通过卷积实现通道数转换和特征深化。
- 功能: U − N e t U-Net U−Net 的核心特征提取单元,通过两次 3 × 3 3×3 3×3 卷积(
-
Down(下采样模块):
- 功能:实现 U − N e t U-Net U−Net 编码器的下采样,缩小特征图尺寸(减半),增加通道数,提取高层语义特征。这些高层语义特征可与 Y O L O YOLO YOLO 模型提取的全局特征互补,提升对息肉整体形态的把握。
- 结构: 2 × 2 2×2 2×2 最大池化(下采样,保留关键特征)+ DoubleConv(特征提取)。
- 优势:最大池化能有效降低特征图分辨率,减少计算量,同时保留局部最大值特征。
-
Up(上采样模块,集成注意力门):
- 功能:实现 U − N e t U-Net U−Net 解码器的上采样,扩大特征图尺寸(加倍),与对应的编码器特征图(经注意力加权)拼接,融合低层细节特征和高层语义特征。可利用 Y O L O YOLO YOLO 检测出的息肉边界框信息,在拼接过程中重点关注框内区域特征,提升融合效率。
- 核心组件:
- 注意力门(
attention):集成前面定义的AttentionGate,对编码器特征图( x 2 x2 x2)进行加权,聚焦关键区域。 - 上采样方式:支持双线性插值(
bilinear=True,默认,计算量小、速度快)和转置卷积(bilinear=False,特征提取能力强但计算量大)。
- 注意力门(
- 前向传播逻辑:
- 注意力加权:先用注意力门处理编码器特征图 x 2 x2 x2。
- 上采样:对解码器特征图 x 1 x1 x1 进行上采样,尺寸加倍。
- 尺寸微调:计算上采样后 x 1 x1 x1 与 x 2 x2 x2 的尺寸差异 ( d i f f X 、 d i f f Y ) (diffX、diffY) (diffX、diffY),通过 p a d d i n g padding padding 调整 x 1 x1 x1 尺寸,确保与 x 2 x2 x2 一致(避免拼接时维度错误)。
- 拼接与卷积:在通道维度(
dim=1)拼接 x 2 x2 x2(注意力加权后的编码器特征)和 x 1 x1 x1(上采样后的解码器特征),通过 DoubleConv 融合特征,输出新的特征图。
-
OutConv(输出卷积模块):
- 功能:将解码器最后一层的特征图映射到分割任务的类别数(此处为 1 1 1,二分类:息肉/背景)。输出的分割结果可与 Y O L O YOLO YOLO 检测的边界框结合,实现“定位+分割”的完整息肉检测流程—— Y O L O YOLO YOLO 快速锁定息肉位置,该模块精准勾勒息肉边界。
- 结构: 1 × 1 1×1 1×1 卷积(无 p a d d i n g padding padding,尺寸不变),将多通道特征图转换为单通道预测图。
- 优势: 1 × 1 1×1 1×1 卷积计算量小,能有效整合多通道特征,输出最终的预测 logits(未经过 sigmoid 激活)。
4.3 轻量化注意力U-Net(AttUNet)整体模型
# 轻量化AttUNet(保持不变)
class AttUNet(nn.Module):
def __init__(self, n_channels=3, n_classes=1, bilinear=True, dropout=0.2):
super(AttUNet, self).__init__()
self.n_channels = n_channels
self.n_classes = n_classes
self.bilinear = bilinear
# 减少通道数,减轻模型复杂度
self.inc = DoubleConv(n_channels, 32, dropout)
self.down1 = Down(32, 64, dropout)
self.down2 = Down(64, 128, dropout)
self.down3 = Down(128, 256, dropout)
factor = 2 if bilinear else 1
self.down4 = Down(256, 512 // factor, dropout)
# 相应调整上采样通道
self.up1 = Up(512, 256 // factor, skip_channels=256, bilinear=bilinear, dropout=dropout)
self.up2 = Up(256, 128 // factor, skip_channels=128, bilinear=bilinear, dropout=dropout)
self.up3 = Up(128, 64 // factor, skip_channels=64, bilinear=bilinear, dropout=dropout)
self.up4 = Up(64, 32, skip_channels=32, bilinear=bilinear, dropout=dropout)
self.outc = OutConv(32, n_classes)
# 初始化权重
self._initialize_weights()
def forward(self, x):
x1 = self.inc(x)
x2 = self.down1(x1)
x3 = self.down2(x2)
x4 = self.down3(x3)
x5 = self.down4(x4)
x = self.up1(x5, x4)
x = self.up2(x, x
六、结果展示
1. 检测任务结果
1.1 量化指标
| 指标 | 数值 | 说明 |
|---|---|---|
| mAP50 | 0.8735 | IoU=0.5 时的平均精度 |
| mAP50-95 | 0.6332 | IoU=0.5-0.95 的平均精度 |
| 精确率(P) | 0.9065 | 预测为息肉的样本准确率 |
| 召回率(R) | 0.7706 | 真实息肉的检出率 |
| F1分数(F) | 0.8330 | 息肉的F1分数结果 |
1.2 可视化效果
- 训练指标记录:

- 训练输出:

- 验证集标签输出(建议与下面的真实预测进行对比查看):

- 验证集预测输出:

位置检测任务总体上准确率相对较高,能够有效检测大致位置,但是碍于样本总数较少,训练效果并为达到最优,小周这里的数据已经经过了数据增强处理,您也可以在此基础上进一步修改,以提高准确率和检测置信度。
2. 分割任务结果
2.1 量化指标
| 指标 | 数值 | 说明 |
|---|---|---|
| Dice 系数 | 0.4405 | 预测与真实掩码的重叠度 |
| IoU(交并比) | 0.2837 | 预测区域与真实区域交集占比 |
| 精确率 | 0.3800 | 预测息肉区域的准确率 |
| 召回率 | 0.5414 | 真实息肉区域的覆盖度 |
2.2 可视化效果
- 训练曲线(
metrics\_curve.png):- 4 个子图:训练 / 验证损失曲线、Dice-IoU 变化曲线、精确率曲线、召回率曲线;
- 收敛特征:通常 50-80 轮达到收敛,早停机制避免过拟合。

-
预测对比图(
predictions.png):- 3 列布局:原始内镜图像、真实灰度掩码(白色为息肉)、预测掩码(含单样本 Dice 值);
- 后处理优化:自动过滤面积 < 50 像素的噪点区域,提升掩码纯净度。

因为样本较少(仅为1000个样本),小周也尝试了模糊,变形等多重的数据增强方式,但是效果都不太理想,使用较大的模型时,发生过拟合以及陷入局部最优,上图的预测就可以很好的展示(模型倾向于将大范围的面积认作是结果,以此达到一个 0.5 Dice系数的局部最优解)。如果您有更好的思路,欢迎进行交流补充,这里小周就是抛砖引玉了。
如果你喜欢我的文章,不妨给小周一个免费的点赞和关注吧!
更多推荐
所有评论(0)