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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 04:35:18