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

如何在tf.keras中不使用sklearn计算macro f1-score

修复tf.keras中Macro F1-Score计算的问题

原代码存在的核心问题

  • 未正确计算Macro F1:f1_macro函数仅保留了最后一个类别的F1值,未对所有类别的F1求平均,导致结果并非真正的Macro F1(Macro F1是所有类别F1值的算术平均值)。
  • 混淆矩阵指标计算冗余:使用列表推导式结合np.multiply的方式效率低,且布尔值与数值相乘易出现类型逻辑错误。
  • 未处理除零异常:当某类别无预测正样本或无真实正样本时,会触发除以零错误,原代码的判断逻辑不完整。
  • 函数逻辑缺失:f1_macro函数未完成核心的平均计算步骤,也无返回值。

修正后的完整代码

import numpy as np
import tensorflow as tf

class Metrics(tf.keras.callbacks.Callback):
    def __init__(self, train_data, val_data):
        self.train_data = train_data
        self.validation_data = val_data
  
    def on_train_begin(self, logs={}):
        self.f1_score_test = []
  
    def f1(self, y_true, y_pred):
        # 用numpy布尔运算高效计算混淆矩阵指标
        TP = np.sum((y_pred == 1) & (y_true == 1))
        FP = np.sum((y_pred == 1) & (y_true == 0))
        FN = np.sum((y_pred == 0) & (y_true == 1))
        
        # 处理分母为0的情况,避免运行时错误
        precision = TP / (TP + FP) if (TP + FP) != 0 else 0.0
        recall = TP / (TP + FN) if (TP + FN) != 0 else 0.0
        
        if precision == 0 or recall == 0:
            return 0.0
        return 2 * (precision * recall) / (precision + recall)

    def f1_macro(self, y_true, y_pred):
        macro_f1 = 0.0
        unique_classes = np.unique(y_true)
        num_classes = len(unique_classes)
        
        for cls in unique_classes:
            # 将当前类别视为正类,其余视为负类
            modified_true = np.where(y_true == cls, 1, 0)
            modified_pred = np.where(y_pred == cls, 1, 0)
            cls_f1 = self.f1(modified_true, modified_pred)
            macro_f1 += cls_f1
        
        # 计算所有类别F1的算术平均值
        return macro_f1 / num_classes if num_classes != 0 else 0.0
  
    def on_epoch_end(self, epoch, logs={}):
        # 获取测试集预测结果并做四舍五入处理
        y_pred_test = np.asarray(self.model.predict(self.validation_data[0])).round()
        # 确保真实标签为一维数组,避免维度不匹配
        y_true_test = np.squeeze(self.validation_data[1])
        
        # 计算并记录Macro F1
        current_macro_f1 = self.f1_macro(y_true_test, y_pred_test)
        self.f1_score_test.append(current_macro_f1)
        
        print(f'Epoch {epoch+1} - f1_test_macro = {current_macro_f1:.4f}')

# 初始化回调实例
new_metrics = Metrics((x_train, y_train), (x_test, y_test))

关键修改说明

  1. 简化混淆矩阵计算:用numpy布尔索引替代低效的列表推导式,逻辑更清晰,计算速度更快。
  2. 完善除零处理:对precision和recall的分母单独判断,避免因无正样本导致的报错。
  3. 实现正确的Macro F1逻辑:遍历每个类别计算F1,累加后除以类别数得到平均,符合Macro F1的定义。
  4. 优化数据格式:用np.squeeze确保真实标签为一维数组,适配多分类/二分类场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 20:39:19