You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何调整损失函数训练Keras中的不平衡多标签分类CNN模型?

解决多标签CNN中加权BCE的pos_weight设置及类别扩展问题

一、如何设置合适的pos_weight避免全0/全1预测

  1. 先计算初始参考值
    不要凭经验硬设数值,先统计训练集的全局正负标签比例:
    假设训练集共有M张图片,每张平均有k个正标签(你的场景是2-4,取平均3),当前训练类别数为C(比如60),那么全局正标签总数是M*k,负标签总数是M*(C - k),初始pos_weight可以设为(C - k)/k。以60类为例,(60-3)/3=19,这是一个合理的起始值,而非直接跳到20、30这类极端值。

  2. 基于验证指标逐步微调
    你已经在跟踪precision和recall,直接用这两个指标调整:

    • 若模型全预测1:说明pos_weight过大,正类权重太高,模型为了降低损失会偏向正类,此时减小pos_weight(比如从19降到15,再逐步试到10);
    • 若模型全预测0:说明pos_weight太小,负类损失占比过高,模型倾向于输出0,此时增大pos_weight;
    • 目标是让precision和recall尽量平衡,或匹配你的业务需求(比如优先召回就稍调高,优先精确就稍调低)。建议尽快实现F1 Score指标,它是precision和recall的调和平均,能更直观反映平衡效果,代码示例:
      import tensorflow as tf
      from tensorflow.keras import backend as K
      
      def f1_score(y_true, y_pred):
          # 将预测值转为0/1(默认用0.5作为阈值,可按需调整)
          y_pred = K.round(y_pred)
          # 计算TP、FP、FN
          tp = K.sum(K.cast(y_true * y_pred, 'float'), axis=0)
          fp = K.sum(K.cast((1 - y_true) * y_pred, 'float'), axis=0)
          fn = K.sum(K.cast(y_true * (1 - y_pred), 'float'), axis=0)
          # 计算精确率和召回率
          p = tp / (tp + fp + K.epsilon())  # 加epsilon避免除0错误
          r = tp / (tp + fn + K.epsilon())
          # 计算F1值
          f1 = 2 * p * r / (p + r + K.epsilon())
          # 处理nan值(当某类无正样本时)
          f1 = tf.where(tf.math.is_nan(f1), tf.zeros_like(f1), f1)
          return K.mean(f1)
      
  3. 考虑逐类加权(而非全局pos_weight)
    如果不同类别的正负样本差异极大(比如部分类正样本极少),全局pos_weight可能不够精准。可以对每个类别单独计算权重:对类别i,统计训练集中该类的负样本数/正样本数,作为该类的权重。然后自定义损失函数,对每个类的BCE损失乘以对应权重:

    def class_weighted_bce(y_true, y_pred, class_weights):
        bce = K.binary_crossentropy(y_true, y_pred)
        # 每个样本的损失乘以对应类的权重
        weighted_bce = y_true * class_weights * bce + (1 - y_true) * bce
        return K.mean(weighted_bce)
    

    其中class_weights是长度等于类别数的数组,每个元素对应该类的权重。

二、扩展到608类时的通用规则

  1. 先做数据分布统计
    不管类别数量多少,第一步必须统计数据的正负样本分布:

    • 全局分布:计算整个数据集所有标签中,负标签总数/正标签总数,作为全局pos_weight的初始值;
    • 逐类分布:对608个类别分别统计该类的负样本数/正样本数,得到每个类的单独权重,适合类别不平衡差异大的场景。
  2. 结合验证指标动态优化
    类别越多,全局不平衡可能越严重(每张图仅2-4个正类,608类的话负类是604-606个,正负比例约1:150-1:300),此时初始pos_weight会很大,但直接使用可能导致模型全1,所以必须从较低的起始值逐步调整,观察验证集的precision、recall、F1,找到平衡点。

  3. 尝试更鲁棒的损失函数
    当类别极多、不平衡极严重时,加权BCE可能不够,可以试试Focal Loss,它通过降低易分类样本的损失权重,让模型更聚焦难分类的正样本:

    def focal_loss(y_true, y_pred, alpha=0.25, gamma=2.0):
        y_pred = K.clip(y_pred, K.epsilon(), 1 - K.epsilon())
        # 计算交叉熵
        cross_entropy = -y_true * K.log(y_pred) - (1 - y_true) * K.log(1 - y_pred)
        # 计算聚焦权重因子
        pt = tf.where(y_true == 1, y_pred, 1 - y_pred)
        focal_weight = alpha * K.pow(1 - pt, gamma)
        return K.mean(focal_weight * cross_entropy)
    

    其中alpha是正负类的平衡系数,gamma是聚焦参数(通常取2),可根据验证指标调整。

  4. 辅助数据策略
    除了损失加权,还可以对正样本做针对性的数据增强(如翻转、裁剪等),或者采用标签平滑,避免模型过于极端地输出0或1。

内容的提问来源于stack exchange,提问作者mistermooster

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 08:40:02