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

使用PredictionErrorDisplay绘制误差图为何报ValueError?

解决PredictionErrorDisplay的ValueError问题

这个错误的核心原因是你传入的test_targets或test_predictions是二维结构(比如形状为(n_samples, 1)的数组/数据框),而PredictionErrorDisplay要求输入必须是一维的(长度等于样本数的序列)。

快速解决步骤:

  1. 先检查输入数据的形状:
print(test_targets.shape, test_predictions.shape)

如果输出是类似(1000, 1)这样的二维结果,就需要转换成一维。

  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()
  1. 修改后的完整代码:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 17:45:03