如何在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))
关键修改说明
- 简化混淆矩阵计算:用numpy布尔索引替代低效的列表推导式,逻辑更清晰,计算速度更快。
- 完善除零处理:对precision和recall的分母单独判断,避免因无正样本导致的报错。
- 实现正确的Macro F1逻辑:遍历每个类别计算F1,累加后除以类别数得到平均,符合Macro F1的定义。
- 优化数据格式:用
np.squeeze确保真实标签为一维数组,适配多分类/二分类场景。
内容的提问来源于stack exchange,提问作者buzz bowlekar
相关产品推荐
相关产品推荐

