使用PredictionErrorDisplay绘制误差图为何报ValueError?
解决PredictionErrorDisplay的ValueError问题
这个错误的核心原因是你传入的test_targets或test_predictions是二维结构(比如形状为(n_samples, 1)的数组/数据框),而PredictionErrorDisplay要求输入必须是一维的(长度等于样本数的序列)。
快速解决步骤:
- 先检查输入数据的形状:
print(test_targets.shape, test_predictions.shape)
如果输出是类似(1000, 1)这样的二维结果,就需要转换成一维。
- 根据数据类型做转换:
- 如果是numpy数组:
import numpy as np test_targets = test_targets.ravel() # 或者用flatten() test_predictions = test_predictions.ravel()
- 如果是pandas DataFrame/Series:
test_targets = test_targets.squeeze() test_predictions = test_predictions.squeeze()
- 修改后的完整代码:
from sklearn.metrics import PredictionErrorDisplay import matplotlib.pyplot as plt # 先转换数据为一维 test_targets = test_targets.ravel() test_predictions = test_predictions.ravel() errors = PredictionErrorDisplay(y_true=test_targets, y_pred=test_predictions) errors.plot() plt.savefig('Output.png') plt.clf()
额外的误差可视化方法
除了用PredictionErrorDisplay,你还可以直接手动绘制误差分布,更直观判断是普遍误差大还是少数异常值拉高标准:
- 误差直方图(看误差分布):
errors = test_predictions - test_targets plt.hist(errors, bins=50, edgecolor='black') plt.xlabel('预测误差') plt.ylabel('样本数量') plt.title('预测误差分布直方图') plt.savefig('error_hist.png') plt.clf()
- 真实值vs预测值散点图(对比偏差):
plt.scatter(test_targets, test_predictions, alpha=0.6) # 画一条完美预测的参考线 plt.plot([test_targets.min(), test_targets.max()], [test_targets.min(), test_targets.max()], 'r--', linewidth=2) plt.xlabel('真实值') plt.ylabel('预测值') plt.title('真实值与预测值对比') plt.savefig('true_vs_pred_scatter.png') plt.clf()
内容的提问来源于stack exchange,提问作者SRJCoding
相关产品推荐
相关产品推荐

