ARTICLE DETAIL

资讯详情

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

PaddleSeg 中的 OhemEdgeAttentionLoss:面向困难样本与边缘提取的在线难例挖掘损失函数

PaddleSeg 中的 OhemEdgeAttentionLoss:面向困难样本与边缘提取的在线难例挖掘损失函数 人工智能计算机视觉预训练【免费下载链接】PaddleSegEasy-to-use image segmentation library with awesome pre-trained model zoo, supporting wide-range of practical tasks in Semantic Segmentation, Interactive Segmentation, Panoptic Segmentation, Image Matting, 3D Segmentation, etc.项目地址https://gitcode.com/gh_mirrors/pa/PaddleSeg点击查看免费下载本文聚焦 PaddleSeg 语义分割工具库paddleseg内置的OhemEdgeAttentionLoss损失函数从设计动机、参数语义、源码级前向计算流程到真实训练配置以 DecoupledSegNet 为例逐层展开。读完本文你将掌握该损失函数的四个核心参数edge_threshold、thresh、min_kept、ignore_index的调优方法理解其边缘筛选 OHEM 困难样本挖掘的双重机制并能在自己的 PaddleSeg 训练配置中正确组合使用它。一、为什么需要 OhemEdgeAttentionLoss困难样本与边缘提取的痛点在语义分割任务中不同像素的学习难度差异巨大类别边界附近、小目标内部、遮挡区域等位置的像素往往分类置信度低、损失大属于困难样本hard samples。如果对所有样本一视同仁地计算交叉熵损失大量易分类像素如大面积的背景区域会稀释困难样本的梯度贡献导致模型在细节、边缘处的分割精度难以提升。OHEMOnline Hard Example Mining在线难例挖掘的核心思想正是针对这一问题根据输入到模型中的各样本产生的损失大小来区分困难样本只保留损失较大即分类精度差的样本参与反向传播从而把模型的注意力集中到真正难学的像素上。在此基础上OhemEdgeAttentionLoss进一步叠加了**边缘注意力Edge Attention**机制先借助边缘预测分支edge logit筛出位于边缘区域的像素再对这部分像素做 OHEM 难例挖掘最终只在这些困难边缘像素上计算交叉熵损失。这使得它在存在困难样本、且希望重点提升边缘提取性能的场景下非常有效——这正是官方文档OhemEdgeAttentionLoss_cn.md给出的典型使用场景。二、类签名与四个核心参数该损失函数在 PaddleSeg 中通过装饰器manager.LOSSES.add_component注册为可配置组件见 manager.py 中的LOSSES ComponentManager(losses)类定义位于 ohem_edge_attention_loss.py。官方 API 签名如下class paddleseg.models.losses.OhemEdgeAttentionLoss( edge_threshold 0.8, thresh 0.7, min_kept 5000, ignore_index 255 )四个参数的含义与默认值如下参数类型默认值作用edge_thresholdfloat0.8边缘判定阈值。预测的边缘概率edge logit大于该值的像素被视作边缘像素只有这些像素会进入后续损失计算。threshfloat0.7OHEM 阈值。像素在真实类别上的预测概率softmax 后取值小于thresh时被视为困难样本予以保留。min_keptint5000最小保留像素数。与thresh配合使用如果thresh设置过高可能导致本轮迭代中没有样本输入损失函数因此该值保证至少前min_kept个概率最低的元素不会被过滤掉。ignore_indexint64255忽略像素值。标注图中标记为该值的像素不参与损失计算对输入梯度不产生贡献常用于标注图中无法标注或很难标注的区域。其中min_kept与thresh的配合逻辑需要特别注意OHEM 并不是简单地固定丢弃概率高于thresh的所有像素而是有一个动态兜底机制——当按概率排序后第min_kept个像素的概率如果仍高于thresh则实际阈值会被提升为该像素的概率详见下文源码解析。这样既能保证保留足够的样本数又能确保参与计算的样本都是当前批次里相对最难的。三、源码级剖析前向计算的两阶段流程OhemEdgeAttentionLoss继承自paddle.nn.Layer其forward接收的是(seg_logit, edge_logit)元组seg_logit是语义分割主分支输出形状为(N, C, H, W)C为类别数edge_logit是边缘预测分支输出C 1形状需与 label 一致。整体流程可拆解为边缘筛选与OHEM 难例挖掘两个阶段实现细节见 ohem_edge_attention_loss.py。阶段一边缘筛选Edge Filtering# Filter out edge filler paddle.ones_like(label) * self.ignore_index label paddle.where(edge_logit self.edge_threshold, label, filler)构造一个全为ignore_index默认 255的填充张量凡edge_logit edge_threshold的位置保留真实标签其余位置全部置为ignore_index。这样非边缘像素在后续交叉熵计算中被自动忽略损失只聚焦于边缘像素。阶段二OHEM 难例挖掘n, c, h, w seg_logit.shape label label.reshape((-1, )) valid_mask (label ! self.ignore_index).astype(int64) num_valid valid_mask.sum() label label * valid_mask prob F.softmax(seg_logit, axis1) prob prob.transpose((1, 0, 2, 3)).reshape((c, -1))先统计有效像素数num_valid再对分割 logit 在类别维做 softmax得到每个像素属于其真实类别的概率。随后进入 OHEM 核心逻辑仅当min_kept num_valid且num_valid 0时执行if self.min_kept num_valid and num_valid 0: prob prob (1 - valid_mask).astype(prob.dtype) # 被忽略像素概率置为 1 label_onehot F.one_hot(label, c) label_onehot label_onehot.transpose((1, 0)) prob prob * label_onehot # 只取真实类别对应的概率 prob paddle.sum(prob, axis0) threshold self.thresh if self.min_kept 0: index prob.argsort() # 按概率升序排序 threshold_index index[min(len(index), self.min_kept) - 1] threshold_index int(threshold_index) if prob[threshold_index] self.thresh: threshold prob[threshold_index] # 动态兜底提升实际阈值 kept_mask (prob threshold).astype(int64) label label * kept_mask valid_mask valid_mask * kept_mask关键点通过 one-hot 提取每个像素在真实类别上的 softmax 概率该概率越低说明模型对它的分类越不确定即越难按概率升序排序后取第min_kept个位置的像素概率作为参考若它仍高于thresh说明当前批次样本整体都较简单实际阈值自动上调为该概率确保至少保留min_kept个像素这正是文档中min_kept保证至少前min_kept个元素不会被过滤掉的代码实现最终保留概率小于实际阈值的像素作为困难样本其余像素的 label 与 valid_mask 一并清零。阶段三掩码化交叉熵与归一化# make the invalid region as ignore label label (1 - valid_mask) * self.ignore_index label label.reshape((n, 1, h, w)) valid_mask valid_mask.reshape((n, 1, h, w)).astype(float32) loss F.softmax_with_cross_entropy(seg_logit, label, ignore_indexself.ignore_index, axis1) loss loss * valid_mask avg_loss paddle.mean(loss) / (paddle.mean(valid_mask) self.EPS)被过滤的像素重新标记为ignore_index使用softmax_with_cross_entropy计算逐像素损失后乘上valid_mask清零无效位置最后除以有效像素均值加EPS 1e-10防除零得到归一化的平均损失。label与valid_mask均设置stop_gradient True仅让分割 logit 接收梯度。对比参考PaddleSeg 中还提供无边缘筛选的通用版本OhemCrossEntropyLoss实现见 ohem_cross_entropy_loss.py其默认min_kept 10000且不支持edge_threshold。两者共享同一套 OHEM 挖掘逻辑OhemEdgeAttentionLoss相当于在它前面加了一道边缘注意力门控。四、在真实训练配置中使用DecoupledSegNet 的组合损失在 PaddleSeg 中损失函数通过配置文件YAML的loss字段组合使用。以 decoupledsegnet_resnet50_os8_cityscapes_832x832_80k.yml 为例这是OhemEdgeAttentionLoss的官方实际落地场景model: type: DecoupledSegNet backbone: type: ResNet50_vd output_stride: 8 multi_grid: [1, 2, 4] pretrained: https://bj.bcebos.com/paddleseg/dygraph/resnet50_vd_ssld_v2.tar.gz num_classes: 19 backbone_indices: [0, 3] aspp_ratios: [1, 12, 24, 36] aspp_out_channels: 256 align_corners: False loss: types: - type: OhemCrossEntropyLoss - type: RelaxBoundaryLoss - type: BCELoss weight: dynamic edge_label: True - type: OhemEdgeAttentionLoss coef: [1, 1, 25, 1] train_dataset: transforms: - type: ResizeStepScaling min_scale_factor: 0.75 max_scale_factor: 2.0 scale_step_size: 0.25 - type: RandomPaddingCrop crop_size: [832, 832] - type: RandomHorizontalFlip - type: RandomDistort brightness_range: 0.4 contrast_range: 0.4 saturation_range: 0.4 - type: Normalize edge: True要点解读多损失组合训练时同时使用OhemCrossEntropyLoss主体分割、RelaxBoundaryLoss边界松弛、BCELoss边缘二分类edge_label: True表示监督 edge 分支与OhemEdgeAttentionLoss边缘难例挖掘由coef: [1, 1, 25, 1]加权求和total_loss 1 * loss1 1 * loss2 25 * loss3 1 * loss4输出对齐DecoupledSegNet在训练模式下返回[seg_logit, body_logit, edge_logit, (seg_logit, edge_logit)]见 decoupled_segnet.py其中最后一个元素正是OhemEdgeAttentionLoss所需的(seg_logit, edge_logit)元组输入与文档中损失类型顺序与模型输出一致的约定吻合数据侧配套train_dataset中开启edge: True表示数据集加载时额外生成边缘标签图供BCELoss与边缘筛选使用。若使用该损失而数据集未提供边缘标注训练将无法正确进行配置校验PaddleSeg 在启动前会通过 config_checker.py 中的DefaultLossRule校验loss配置必须同时包含types与coef且两者长度一致当types长度为 1 时会自动复制到与coef等长同时DefaultSyncIgnoreIndexRule会同步各损失的ignore_index。五、参数调优指南与使用建议场景 / 现象建议调整边缘预测分支置信度普遍偏低边缘区域被大量过滤适当降低edge_threshold如 0.7让更多像素进入候选集边缘预测分支过强、噪声边缘混入适当提高edge_threshold如 0.9收紧边缘判定训练中损失出现空批次或震荡剧烈提高min_kept保证每轮都有足量样本参与损失计算希望挖掘更难样本、加大梯度集中度降低thresh如 0.6只保留更难的低置信度像素标注图存在大量未标注/难标注区域将ignore_index设置为该区域的标注灰度值默认 255需要强调的是thresh与min_kept并非独立生效实际阈值 max(thresh, 排序后第 min_kept 个像素的概率)因此调参时应二者联动。另外该损失只处理分割主分支的交叉熵部分不直接约束边缘分支本身——边缘分支的监督由BCELossedge_label: True负责二者各司其职共同构成 DecoupledSegNet 的主体 边缘 难例挖掘训练体系。六、直接使用 API 的调用示例在脱离配置文件、直接以 Python API 使用 PaddleSeg 的场景下可以这样实例化并调用该损失import paddle from paddleseg.models.losses import OhemEdgeAttentionLoss criterion OhemEdgeAttentionLoss( edge_threshold0.8, # 边缘判定阈值 thresh0.7, # OHEM 阈值 min_kept5000, # 最小保留像素数 ignore_index255, # 忽略像素值 ) # logits: (seg_logit, edge_logit)形状分别为 (N, C, H, W) 与 (N, 1, H, W) # label: (N, H, W)dtype 为 int64 seg_logit paddle.randn([2, 19, 64, 64]) edge_logit paddle.sigmoid(paddle.randn([2, 1, 64, 64])) label paddle.randint(0, 19, [2, 64, 64]).astype(int64) loss criterion((seg_logit, edge_logit), label) print(loss.item())注意forward的第一个入参必须是(seg_logit, edge_logit)元组代码中通过logits[0], logits[1]解包且edge_logit的形状必须与label一致否则会抛出ValueError见 ohem_edge_attention_loss.py。七、小结OhemEdgeAttentionLoss将边缘注意力筛选与OHEM 在线难例挖掘有机融合先以edge_threshold把损失范围收窄到边缘像素再以threshmin_kept动态兜底机制挑出其中模型最没把握的困难样本最后只在掩码化的困难边缘像素上计算归一化交叉熵。从源码实现ohem_edge_attention_loss.py到真实配置decoupledsegnet_resnet50_os8_cityscapes_832x832_80k.yml它是 PaddleSeg 在 DecoupledSegNet 这类主体—边缘解耦模型上提升边缘分割精度的关键组件之一适用于任何存在困难样本、且需要强化边缘提取性能的分割任务。赞分享人工智能计算机视觉预训练【免费下载链接】PaddleSegEasy-to-use image segmentation library with awesome pre-trained model zoo, supporting wide-range of practical tasks in Semantic Segmentation, Interactive Segmentation, Panoptic Segmentation, Image Matting, 3D Segmentation, etc.项目地址https://gitcode.com/gh_mirrors/pa/PaddleSeg点击查看免费下载相关推荐边界损失函数3步解决图像分割中的边缘精度难题边界损失函数3步解决图像分割中的边缘精度难题 还在为医学影像分割的模糊边界而烦恼吗传统方法在处理高度不平衡数据时往往力不从心特别是当目标区域仅占图像极小比人工智能深度学习计算机视觉医疗健康Effect 中的 Cron 字段校验与星期归一化Cron.make 如何将 weekday 7 视为周日Effect 中的 Cron 字段校验与星期归一化 Cron.make 如何将 weekday 7 视为周日 导读 本文围绕 Effect 库中 Cron 模人工智能计算机视觉预训练PyTorch Metric Learning挖掘器指南如何高效选择困难样本对PyTorch Metric Learning挖掘器指南如何高效选择困难样本对 PyTorch Metric Learning是一个强大的深度学习度量学习库机器学习深度学习上一篇JWT Decode 开源项目教程下一篇pandas 稀疏数据结构SparseArray / SparseDtype完整指南压缩存储、density 度量与 scipy.sparse 互操作创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表