基于 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 数据处理细节

  • 检测任务

    1. 校验图像有效性(排除空图、损坏文件);
    2. CSV 标注(class_name, xmin, ymin, xmax, ymax) 转换为 YOLO 格式(归一化中心坐标 + 宽高)
    3. 按 8:2 比例随机划分训练集 / 验证集,生成 data.yaml 配置文件。
      .
  • 分割任务

    1. 图像与掩码尺寸统一 resize 至 256 × 256 256×256 256×256
    2. 85 : 15 85:15 85:15 比例划分训练集 / 验证集,小样本时自动重复数据增强训练;
    3. 内置多维度数据增强(翻转、旋转、缩放、弹性变换等)。

四、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} IRH×W×3,首先将其 Resize 至固定尺寸 I r e s i z e ∈ R 640 × 640 × 3 I_{resize} \in \mathbb{R}^{640 \times 640 \times 3} IresizeR640×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 0255 的范围缩小到 0 − 1 0-1 01,就像把测量单位从“厘米”换成“米”,让模型计算更高效。

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} FinRC×H×W 分为两部分,一部分直接通过,另一部分经过多组瓶颈层(Bottleneck)处理后与前一部分融合,输出特征图 F o u t ∈ R 2 C × H × W F_{out} \in \mathbb{R}^{2C \times H \times W} FoutR2C×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} F1RC1×H1×W1 F 2 ∈ R C 2 × H 2 × W 2 F_2 \in \mathbb{R}^{C2 \times H2 \times W2} F2RC2×H2×W2 F 3 ∈ R C 3 × H 3 × W 3 F_3 \in \mathbb{R}^{C3 \times H3 \times W3} F3RC3×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} P1R(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} P2R(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} P3R(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=1IoU+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 网络架构图:

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} xRCx×H×W(输入特征)和解码端特征 g ∈ R C g × H g × W g g \in \mathbb{R}^{C_g \times H_g \times W_g} gRCg×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=1X+Y2×XY

其中 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()
    }
  1. 适用场景:用于评估胃息肉分割模型的预测效果,核心指标覆盖分割任务的关键评价维度。
  2. 输入要求:pred 为模型输出经过 sigmoid 激活后的概率图(值在 0 − 1 0-1 01 之间),target 为真实掩码( 0 / 1 0/1 0/1 二值图)。
  3. 核心逻辑:
    • 二值化:将预测概率图按阈值(默认 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 01,越接近 1 1 1 表示重叠度越高。
      • I o U IoU IoU(交并比):预测区域与真实区域的交集除以并集,直观反映分割区域的匹配程度。
  4. 数值稳定性:所有除法都添加 1 e − 8 1e-8 1e8,避免分母为 0 0 0 导致的计算错误。
  5. 输出:返回包含 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)
  1. 设计思路:基于 Dice系数 的损失函数,Dice系数越接近 1 表示分割效果越好,因此损失函数定义为 1 − D i c e 1 - Dice 1Dice 系数,使损失越小分割效果越好。
  2. 关键步骤:
    • 激活函数:输入input为模型未经过 sigmoid 的原始输出,需先通过 sigmoid 转换为 0 − 1 0-1 01 的概率值。
    • 交集计算:通过张量点乘input * target得到预测与真实掩码的交集,求和得到交集大小。
    • 损失计算:引入smooth(默认 1 e − 5 1e-5 1e5)避免分子/分母为 0 0 0,提升数值稳定性。
  3. 相比交叉熵损失,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
  1. 设计思路:融合 二元交叉熵损失(BCEWithLogitsLoss)Dice损失 的优势,兼顾 类别概率分布 的拟合(交叉熵)和 分割区域的重叠度(Dice),提升模型的分割性能。
  2. 组件说明:
    • bce_loss:使用 BCEWithLogitsLoss,直接接收模型原始输出(无需提前 sigmoid),计算预测概率与真实标签的交叉熵,擅长拟合类别分布。
    • dice_loss:调用自定义的 DiceLoss,关注区域重叠度。
  3. 权重调节:通过 bce_weightdice_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$,更侧重区域重叠度。
  4. 解决单一损失的局限性,交叉熵保证模型的梯度稳定性, 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
  1. 功能:模拟图像拍摄时的模糊场景,增加数据多样性,提升模型对模糊图像的鲁棒性。
  2. 参数说明:
    • p:执行该增强的概率(默认 0.5 0.5 0.5),即 50 % 50\% 50% 的图像会被模糊处理。
    • radius_range:高斯模糊核半径的范围(默认 0.5 − 2.0 0.5-2.0 0.52.0),半径越大模糊程度越高。
  3. 执行逻辑:调用时生成随机数,若小于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
  1. 功能:模拟组织形变(如胃壁蠕动导致的息肉形态变化),增强模型对息肉不同形态的适应能力,核心修复了原始弹性变换中可能出现的维度不匹配问题。
  2. 参数说明:
    • p:执行概率(默认 0.4 0.4 0.4)。
    • alpha:控制位移幅度(默认 120 120 120),值越大形变越剧烈。
    • sigma:控制位移场的平滑度(默认 15 15 15),值越大形变越连贯。
  3. 核心逻辑(修复维度不匹配的关键):
    • 先将 PIL图像 转换为 numpy数组,提取图像的 高度( h h h宽度( w w w,忽略通道数,确保位移场尺寸与图像尺寸严格匹配。
    • 生成 2D位移场(dx、dy):通过随机生成 [ − 1 , 1 ] [-1,1] [1,1] 的矩阵,经高斯模糊平滑后乘以 alpha 得到最终位移,保证位移场是与图像尺寸一致的 2D 矩阵。
    • 坐标网格与重映射:创建像素坐标网格 ( x 、 y ) (x、y) xy,叠加位移场得到新坐标 ( m a p _ x 、 m a p _ y ) (map\_x、map\_y) map_xmap_y;通过 cv2.remap 对每个通道(彩色图像)或单通道(灰度图像)进行 重映射,实现弹性形变。
  4. 边界处理:使用 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
  1. 功能:模拟不同拍摄距离下的息肉大小变化,增强模型对息肉不同尺度的识别能力,确保缩放后图像尺寸与原始一致(便于批量训练)。
  2. 参数说明:
    • p:执行概率(默认 0.5 0.5 0.5)。
    • zoom_range:缩放比例范围(默认 0.8 − 1.2 0.8-1.2 0.81.2),即图像可缩小至 80 % 80\% 80% 或放大至 120 % 120\% 120%
  3. 执行逻辑:
    • 随机选择缩放比例,先将图像 resize 到新尺寸(使用 双线性插值 Image.BILINEAR,保证图像质量)。
    • 缩放后裁剪回原尺寸:
      • 缩小(zoom<1.0):新尺寸大于原尺寸,随机选择裁剪区域(确保裁剪后尺寸与原尺寸一致),增加随机性。
      • 放大(zoom>1.0):新尺寸大于原尺寸,中心裁剪(保证息肉主体不被裁掉)。
  4. 避免缩放后尺寸不一致导致的批量训练错误,同时通过随机/中心裁剪平衡随机性和主体保留需求。

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
  1. 类定位:继承 P y T o r c h PyTorch PyTorchDataset 类,自定义适用于胃息肉分割任务的数据集,实现图像与对应掩码的加载、增强和格式转换。
  2. __init__方法(初始化):
    • 输入参数:图像路径列表(image_paths)、掩码路径列表(mask_paths)、变换字典(transform)、是否为训练集(is_train)。
    • 重复因子(repeat):针对小样本场景(训练集样本数< 1500 1500 1500),将训练集重复 2 2 2 次,增加训练迭代次数,提升模型泛化能力;验证集不重复。
  3. __len__方法(数据集长度):返回样本数乘以重复因子,决定训练时的迭代次数。
  4. __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.71.3)和亮度增强( 30 % 30\% 30% 概率,亮度因子 0.7 − 1.3 0.7-1.3 0.71.3),进一步丰富训练数据多样性。
    • 变换应用:通过transform字典分别对图像(train变换链)和掩码(mask变换链)进行处理(如resize、随机翻转、归一化等)。
    • 掩码二值化:将掩码转换为 0 / 1 0/1 0/1 二值张量(阈值 0.5 0.5 0.5),符合分割任务的标签要求。
  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
  1. 核心作用:作为 U − N e t U-Net UNet 解码器与编码器特征融合的桥梁,通过门控信号(解码器特征)引导编码器特征的注意力加权,使模型聚焦于息肉区域的关键特征,抑制背景噪声,提升分割精度。该模块与 Y O L O YOLO YOLO 检测模型形成互补—— Y O L O YOLO YOLO 负责快速定位息肉大致区域(边界框),注意力模块则助力分割模型精准捕捉息肉边缘细节。
  2. __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 01 的注意力权重图。
    • 辅助组件: R e L U ReLU ReLU 激活函数(引入非线性)、上采样层(用于尺寸匹配)。
  3. forward 方法(前向传播)
    • 输入:x(编码器特征图)、g(解码器特征图,门控信号)。
    • 特征映射:通过 gatetransform 分支分别处理门控信号和编码器特征。
    • 尺寸匹配(关键改进):若两者尺寸不一致,通过上采样(门控信号尺寸小时)或自适应最大池化(编码器特征尺寸小时)调整,确保后续能直接相加,解决 U − N e t U-Net UNet 不同层级特征尺寸不匹配的问题。
    • 注意力权重生成:将匹配后的特征相加,经 R e L U ReLU ReLU 激活后通过 psi分支 生成注意力权重图。
    • 权重应用:将编码器特征与注意力权重图逐元素相乘,强化关键区域特征,抑制无关背景特征。此过程可结合 Y O L O YOLO YOLO 检测出的息肉边界框区域,进一步聚焦有效区域,减少无效背景特征的干扰。
  4. 输出:经过注意力加权后的编码器特征,用于后续与解码器特征拼接融合。

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)
  1. DoubleConv(双重卷积模块)

    • 功能: U − N e t U-Net UNet 的核心特征提取单元,通过两次 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,通过卷积实现通道数转换和特征深化。
  2. Down(下采样模块)

    • 功能:实现 U − N e t U-Net UNet 编码器的下采样,缩小特征图尺寸(减半),增加通道数,提取高层语义特征。这些高层语义特征可与 Y O L O YOLO YOLO 模型提取的全局特征互补,提升对息肉整体形态的把握。
    • 结构: 2 × 2 2×2 2×2 最大池化(下采样,保留关键特征)+ DoubleConv(特征提取)。
    • 优势:最大池化能有效降低特征图分辨率,减少计算量,同时保留局部最大值特征。
  3. Up(上采样模块,集成注意力门)

    • 功能:实现 U − N e t U-Net UNet 解码器的上采样,扩大特征图尺寸(加倍),与对应的编码器特征图(经注意力加权)拼接,融合低层细节特征和高层语义特征。可利用 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) diffXdiffY,通过 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 融合特征,输出新的特征图。
  4. 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 可视化效果

  1. 训练曲线metrics\_curve.png):
    • 4 个子图:训练 / 验证损失曲线、Dice-IoU 变化曲线、精确率曲线、召回率曲线;
    • 收敛特征:通常 50-80 轮达到收敛,早停机制避免过拟合。

请添加图片描述

  1. 预测对比图predictions.png):

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

请添加图片描述
因为样本较少(仅为1000个样本),小周也尝试了模糊,变形等多重的数据增强方式,但是效果都不太理想,使用较大的模型时,发生过拟合以及陷入局部最优,上图的预测就可以很好的展示(模型倾向于将大范围的面积认作是结果,以此达到一个 0.5 Dice系数的局部最优解)。如果您有更好的思路,欢迎进行交流补充,这里小周就是抛砖引玉了。

如果你喜欢我的文章,不妨给小周一个免费的点赞和关注吧!

Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐