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

TensorFlow自定义训练循环中二进制分类任务的F1-score评估

TensorFlow自定义训练循环中二进制分类的F1-Score实现

要在自定义训练循环中像使用BinaryAccuracy、AUC一样集成F1-score评估,直接继承tf.keras.metrics.Metric类实现即可。核心思路是跟踪真阳性(TP)、假阳性(FP)、假阴性(FN)三个关键指标,再通过这三个值计算F1-score。

自定义BinaryF1Score类实现

import tensorflow as tf

class BinaryF1Score(tf.keras.metrics.Metric):
    def __init__(self, threshold=0.5, name='binary_f1_score', **kwargs):
        super().__init__(name=name, **kwargs)
        self.threshold = threshold
        # 初始化累计状态变量
        self.true_positives = self.add_weight(name='tp', initializer='zeros')
        self.false_positives = self.add_weight(name='fp', initializer='zeros')
        self.false_negatives = self.add_weight(name='fn', initializer='zeros')

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 将预测值按阈值转为二分类标签
        y_pred = tf.cast(tf.greater(y_pred, self.threshold), tf.float32)
        y_true = tf.cast(y_true, tf.float32)

        # 计算当前batch的TP、FP、FN
        tp = tf.reduce_sum(y_true * y_pred)
        fp = tf.reduce_sum((1 - y_true) * y_pred)
        fn = tf.reduce_sum(y_true * (1 - y_pred))

        # 累加到全局状态
        self.true_positives.assign_add(tp)
        self.false_positives.assign_add(fp)
        self.false_negatives.assign_add(fn)

        # 处理样本权重(可选)
        if sample_weight is not None:
            sample_weight = tf.cast(sample_weight, tf.float32)
            tp_weighted = tf.reduce_sum(sample_weight * y_true * y_pred)
            fp_weighted = tf.reduce_sum(sample_weight * (1 - y_true) * y_pred)
            fn_weighted = tf.reduce_sum(sample_weight * y_true * (1 - y_pred))
            self.true_positives.assign_add(tp_weighted)
            self.false_positives.assign_add(fp_weighted)
            self.false_negatives.assign_add(fn_weighted)

    def result(self):
        # 计算精确率和召回率,用divide_no_nan避免除以0
        precision = tf.math.divide_no_nan(self.true_positives, self.true_positives + self.false_positives)
        recall = tf.math.divide_no_nan(self.true_positives, self.true_positives + self.false_negatives)
        # 计算F1-score
        return tf.math.divide_no_nan(2 * precision * recall, precision + recall)

    def reset_state(self):
        # 重置状态变量,为下一轮评估做准备
        self.true_positives.assign(0.0)
        self.false_positives.assign(0.0)
        self.false_negatives.assign(0.0)

在自定义训练循环中的使用示例

和官方指标的调用方式完全一致,直接在训练、验证步骤中调用update_state、result、reset_state方法:

# 初始化指标
train_acc_metric = tf.keras.metrics.BinaryAccuracy()
train_f1_metric = BinaryF1Score(threshold=0.5)

val_acc_metric = tf.keras.metrics.BinaryAccuracy()
val_f1_metric = BinaryF1Score(threshold=0.5)

# 训练循环
epochs = 10
for epoch in range(epochs):
    print(f"Epoch {epoch+1}/{epochs}")
    # 训练阶段
    for x_batch_train, y_batch_train in train_dataset:
        with tf.GradientTape() as tape:
            y_pred = model(x_batch_train, training=True)
            loss = loss_fn(y_batch_train, y_pred)
        grads = tape.gradient(loss, model.trainable_weights)
        optimizer.apply_gradients(zip(grads, model.trainable_weights))
        
        # 更新训练指标
        train_acc_metric.update_state(y_batch_train, y_pred)
        train_f1_metric.update_state(y_batch_train, y_pred)
    
    # 获取训练指标结果
    train_acc = train_acc_metric.result()
    train_f1 = train_f1_metric.result()
    print(f"训练准确率: {train_acc:.4f}, 训练F1-Score: {train_f1:.4f}")
    
    # 重置训练指标
    train_acc_metric.reset_state()
    train_f1_metric.reset_state()
    
    # 验证阶段
    for x_batch_val, y_batch_val in val_dataset:
        y_pred_val = model(x_batch_val, training=False)
        val_acc_metric.update_state(y_batch_val, y_pred_val)
        val_f1_metric.update_state(y_batch_val, y_pred_val)
    
    # 获取验证指标结果
    val_acc = val_acc_metric.result()
    val_f1 = val_f1_metric.result()
    print(f"验证准确率: {val_acc:.4f}, 验证F1-Score: {val_f1:.4f}\n")
    
    # 重置验证指标
    val_acc_metric.reset_state()
    val_f1_metric.reset_state()

关键说明

  • __init__:初始化分类阈值和三个累计状态变量,用于存储整个epoch的TP、FP、FN
  • update_state:将模型输出转为二分类标签,计算当前batch的指标并累加到全局状态,支持样本权重
  • result:基于累计的TP、FP、FN计算精确率、召回率,最终得到F1-score,用divide_no_nan避免除以0的异常
  • reset_state:每个epoch结束后重置状态,确保下一轮评估从0开始累计

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 10:15:35