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

如何实现并展示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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 16:36:03