如何实现并展示BERT文本分类任务的各项评估指标结果
实现代码说明
你需要的多标签/多分类评估指标可以通过scikit-learn的工具快速实现,以下是完整可复用的代码:
1. 导入依赖
import numpy as np from sklearn.metrics import classification_report, confusion_matrix import matplotlib.pyplot as plt import seaborn as sns
2. 预处理预测结果
首先把模型输出的概率值转成0/1标签(多标签场景),如果是单标签场景直接取概率最高的类别索引即可:
# 多标签场景 # 替换为你自己的模型预测输出和真实标签 y_true = 测试集真实标签 # 形状为(样本数, 类别数),元素为0/1 y_pred_proba = model.predict(测试集) # 阈值可根据业务需求调整,默认用0.5 y_pred = (y_pred_proba >= 0.5).astype(int) # 单标签多分类场景写法参考 # y_pred = np.argmax(y_pred_proba, axis=1)
3. 生成分类评估报告
以下代码会直接输出你截图中展示的每个类别精确率、召回率、F1值,以及微平均、宏平均、加权平均指标:
# 替换为你自己的类别名称列表 category_names = ["类别A", "类别B", "类别C", "其他类别"] report = classification_report( y_true, y_pred, target_names=category_names, digits=4 # 保留4位小数,可调整 ) print(report)
可选:生成混淆矩阵热力图
如果需要可视化混淆矩阵,可以用以下代码:
# 单标签场景 plt.figure(figsize=(8, 6)) cm = confusion_matrix(y_true, y_pred) sns.heatmap( cm, annot=True, fmt='d', cmap='Blues', xticklabels=category_names, yticklabels=category_names ) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.title('分类混淆矩阵') plt.show() # 多标签场景可以按每个类别单独生成混淆矩阵 # for idx, cate in enumerate(category_names): # plt.figure(figsize=(5, 4)) # cm = confusion_matrix(y_true[:, idx], y_pred[:, idx]) # sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['负例', '正例'], yticklabels=['负例', '正例']) # plt.xlabel('预测结果') # plt.ylabel('真实结果') # plt.title(f'{cate} 混淆矩阵') # plt.show()
补充说明
- 微平均:把所有类别的样本合并后计算指标,适合样本不均衡场景
- 宏平均:每个类别的指标直接算数平均,对小类别更敏感
- 加权平均:按每个类别的样本数加权计算平均,兼顾样本量差异
内容的提问来源于stack exchange,提问作者KuugyR
相关产品推荐
相关产品推荐

