如何调整损失函数训练Keras中的不平衡多标签分类CNN模型?
一、如何设置合适的pos_weight避免全0/全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这类极端值。基于验证指标逐步微调
你已经在跟踪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)
考虑逐类加权(而非全局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类时的通用规则
先做数据分布统计
不管类别数量多少,第一步必须统计数据的正负样本分布:- 全局分布:计算整个数据集所有标签中,负标签总数/正标签总数,作为全局pos_weight的初始值;
- 逐类分布:对608个类别分别统计该类的负样本数/正样本数,得到每个类的单独权重,适合类别不平衡差异大的场景。
结合验证指标动态优化
类别越多,全局不平衡可能越严重(每张图仅2-4个正类,608类的话负类是604-606个,正负比例约1:150-1:300),此时初始pos_weight会很大,但直接使用可能导致模型全1,所以必须从较低的起始值逐步调整,观察验证集的precision、recall、F1,找到平衡点。尝试更鲁棒的损失函数
当类别极多、不平衡极严重时,加权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),可根据验证指标调整。辅助数据策略
除了损失加权,还可以对正样本做针对性的数据增强(如翻转、裁剪等),或者采用标签平滑,避免模型过于极端地输出0或1。
内容的提问来源于stack exchange,提问作者mistermooster

