使用TensorFlow API实现带类别权重的Focal Loss报错排查
问题解决:TensorFlow BinaryFocalCrossentropy 参数错误及带类别权重的Focal Loss实现
报错原因
TensorFlow官方的BinaryFocalCrossentropy类没有apply_class_balancing这个初始化参数,你传入了不存在的参数,导致抛出TypeError。
修正步骤及实现方案
1. 修正基础的BinaryFocalCrossentropy使用
你的输出层用了sigmoid激活函数,所以必须将from_logits设为False(from_logits=True仅适用于未经过激活的原始输出)。修正后的loss初始化代码:
tf.keras.losses.BinaryFocalCrossentropy(gamma=2, from_logits=False)
2. 实现带类别权重的Focal Loss
针对不平衡数据,有两种常用方式添加类别权重:
方式一:在训练时传入class_weight参数
先根据训练集标签计算类别权重(正类和负类的权重反比于样本数量):
# 假设y_train是你的训练标签数组 pos_count = np.sum(y_train == 1) neg_count = np.sum(y_train == 0) # 计算权重:让样本少的类别拥有更高权重 class_weight = { 0: pos_count / (pos_count + neg_count), 1: neg_count / (pos_count + neg_count) }
然后在model.fit()中传入该参数:
model.fit( x_train, y_train, class_weight=class_weight, epochs=..., batch_size=..., # 其他训练参数 )
方式二:自定义整合类别权重的Focal Loss
如果希望直接将权重整合到损失函数中,可以自定义实现:
def weighted_binary_focal_loss(gamma=2., alpha=None): def loss(y_true, y_pred): y_true = tf.cast(y_true, tf.float32) # 避免计算log(0)导致数值不稳定 y_pred = tf.clip_by_value(y_pred, 1e-7, 1. - 1e-7) # 计算基础交叉熵 cross_entropy = -y_true * tf.math.log(y_pred) - (1 - y_true) * tf.math.log(1 - y_pred) # 计算Focal Loss的权重项 focal_weight = alpha * y_true * tf.math.pow(1 - y_pred, gamma) + (1 - alpha) * (1 - y_true) * tf.math.pow(y_pred, gamma) # 加权后的Focal Loss focal_loss = focal_weight * cross_entropy return tf.reduce_mean(focal_loss) return loss
使用时,先计算alpha值(通常设为负类样本数/总样本数,或者根据需求调整):
alpha = neg_count / (pos_count + neg_count) # 在compile中使用自定义损失 model.compile( loss=weighted_binary_focal_loss(gamma=2, alpha=alpha), metrics=[tf.keras.metrics.AUC(name='auc'), tf.keras.metrics.Recall()], optimizer=adam )
内容的提问来源于stack exchange,提问作者ForestGump
相关产品推荐
相关产品推荐

