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

Scikit-learn平均精确率评分输入形状异常及PR曲线绘制求助

解决Scikit-learn Average Precision Score输入形状错误问题

嘿,我之前也踩过这个坑!咱们来一步步搞清楚问题出在哪,怎么解决。

问题根源

你用clf.predict_proba(test_set)得到的y_score是一个二维数组,形状是(样本数量, 类别数量)——每个样本对应所有类别的预测概率。但average_precision_score在默认的二分类场景下,需要的是一维数组(只保留正类的预测概率);如果是多分类/多标签任务,也需要调整输入格式才能匹配。

分场景解决方案

情况1:二分类任务(最常见)

这时候你只需要从predict_proba的结果中提取正类的概率列即可。假设LabelEncoder把你的正类编码成了1,代码修改如下:

import matplotlib.pyplot as plt
from sklearn import preprocessing
from sklearn.metrics import average_precision_score, precision_recall_curve

lbl_enc = preprocessing.LabelEncoder()
labels = lbl_enc.fit_transform(test_tags)

# 只保留正类的概率得分([:, 1]表示取第二列,对应编码为1的类别)
y_score = clf.predict_proba(test_set)[:, 1]

average_precision = average_precision_score(labels, y_score)
print('Average precision-recall score: {0:0.2f}'.format(average_precision))

precision, recall, _ = precision_recall_curve(labels, y_score)
plt.step(recall, precision, color='b', alpha=0.2, where='post')
plt.fill_between(recall, precision, alpha=0.2, color='b')
plt.xlabel('Recall')
plt.ylabel('Precision')
plt.title(f'Precision-Recall Curve (AP={average_precision:.2f})')
plt.show()

情况2:多分类任务

如果你的任务是多分类,需要先把标签二值化,同时指定average参数来计算平均精确率:

import matplotlib.pyplot as plt
from sklearn import preprocessing
from sklearn.metrics import average_precision_score, precision_recall_curve

lbl_enc = preprocessing.LabelEncoder()
labels = lbl_enc.fit_transform(test_tags)
# 将标签二值化,每个类别对应一列
labels_binarized = preprocessing.label_binarize(labels, classes=lbl_enc.classes_)
n_classes = labels_binarized.shape[1]

y_score = clf.predict_proba(test_set)
# 选择合适的平均方式,比如'macro'(对每个类别单独计算后取平均)
average_precision = average_precision_score(labels_binarized, y_score, average='macro')
print(f'Average precision-recall score (macro): {average_precision:.2f}')

# 绘制每个类别的PR曲线
plt.figure()
for i in range(n_classes):
    precision, recall, _ = precision_recall_curve(labels_binarized[:, i], y_score[:, i])
    plt.step(recall, precision, alpha=0.2, where='post',
             label=f'Class {lbl_enc.classes_[i]} (AP={average_precision_score(labels_binarized[:, i], y_score[:, i]):.2f})')

plt.fill_between(recall, precision, alpha=0.2, color='b')
plt.xlabel('Recall')
plt.ylabel('Precision')
plt.ylim([0.0, 1.05])
plt.xlim([0.0, 1.0])
plt.title('Multi-class Precision-Recall Curve')
plt.legend(loc="lower left")
plt.show()

情况3:多标签任务

如果你的test_tags是多标签(每个样本可以属于多个类别),那你需要用MultiLabelBinarizer代替LabelEncoder,确保标签和y_score的形状匹配:

import matplotlib.pyplot as plt
from sklearn.preprocessing import MultiLabelBinarizer
from sklearn.metrics import average_precision_score, precision_recall_curve

mlb = MultiLabelBinarizer()
labels = mlb.fit_transform(test_tags)

y_score = clf.predict_proba(test_set)
average_precision = average_precision_score(labels, y_score, average='micro')
print(f'Average precision-recall score (micro): {average_precision:.2f}')

# 绘制PR曲线
precision, recall, _ = precision_recall_curve(labels.ravel(), y_score.ravel())
plt.step(recall, precision, color='b', alpha=0.2, where='post')
plt.fill_between(recall, precision, alpha=0.2, color='b')
plt.xlabel('Recall')
plt.ylabel('Precision')
plt.title(f'Multi-label Precision-Recall Curve (AP={average_precision:.2f})')
plt.show()

关键提醒

  • 一定要确认你的任务类型(二分类/多分类/多标签),不同场景下输入格式要求不同
  • precision_recall_curve的输入要求和average_precision_score一致,所以修改y_score的格式后,曲线绘制部分也能正常运行

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 10:10:05