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

如何获取GridSearchCV最优估计器(best_estimator)的混淆矩阵

获取GridSearchCV最优估计器的混淆矩阵

没问题,其实GridSearchCV在你设置refit='Accuracy'之后,已经自动把最优参数的模型拟合到整个训练集了,并且存在best_estimator_属性里——可能你误以为这个属性没保存?咱们直接用它来生成混淆矩阵就行,步骤很简单:

步骤1:提取最优模型

首先从训练好的GridSearchCV对象里取出最优估计器:

best_rf = gs.best_estimator_

这里best_estimator_会返回用最优参数组合训练完成的RandomForestClassifier实例,因为你设置了refit='Accuracy',它会自动用整个训练数据(X_Distances, Y)拟合这个最优模型。

步骤2:生成预测结果

接下来用这个最优模型生成预测值。如果你的数据已经拆分了训练集和测试集,强烈推荐用测试集来评估(避免过拟合的偏差),比如:

# 假设你有测试集X_test和Y_test
y_pred = best_rf.predict(X_test)

如果暂时没有拆分测试集,也可以用训练集生成预测(但注意这是训练集上的结果,不能代表泛化能力):

y_pred = best_rf.predict(X_Distances)

步骤3:计算并输出混淆矩阵

用sklearn的confusion_matrix工具来计算,还可以可视化让结果更直观:

from sklearn.metrics import confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

# 计算混淆矩阵
cm = confusion_matrix(Y_test, y_pred)  # 用训练集的话替换成Y, y_pred

# 可视化混淆矩阵(可选)
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues')
plt.xlabel('Predicted Label')
plt.ylabel('True Label')
plt.title('Confusion Matrix of Best Random Forest Model')
plt.show()

补充小提示

  • 如果你想确认最优参数组合,可以用print(gs.best_params_)查看,能验证是不是你预期的参数搭配。
  • 因为你设置了cv=3,GridSearchCV会在交叉验证中遍历所有参数组合,选出准确率最高的那个,再用全量训练数据拟合——这就是refit参数的核心作用,所以best_estimator_是完全可用的。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 02:28:09