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

多分类模型预测结果绘图遇ValueError报错,求修复方案

多分类模型预测结果绘图ValueError问题解决

问题场景

需要绘制多分类模型的预测结果,运行以下代码时始终触发ValueError,数据集为企鹅数据集:

def plot_predictions(train_data=X_tr,
                     train_labels=y_tr,
                     test_data=X_t,
                     test_labels=y_t,
                     predictions=None):
  
    plt.figure(figsize=(10, 7))


    plt.scatter(train_data[:,0], train_labels[:,0], c="b", s=4, label="Training data")

    plt.scatter(test_data, test_labels, c="g", s=4, label="Testing data")

    if predictions is not None:
        plt.scatter(test_data, predictions, c="r", s=4, label="Predictions")
        
        correct = (predictions == test_labels).sum().item()
        total = test_labels.size(0)
        accuracy = correct / total
        plt.title(f"Accuracy: {accuracy:.2f}")

    # Show the legend
    plt.legend(prop={"size": 14})
    plt.show()

报错信息

ValueError                                Traceback (most recent call last)
Cell In[1070], line 1
----> 1 plot_predictions();

Cell In[1069], line 10
      7 plt.figure(figsize=(10, 7))
      9 # Plot training data in blue
---> 10 plt.scatter(train_data[:,0], train_labels[:,0], c="b", s=4, label="Training data")
     12 # Plot test data in green
     13 plt.scatter(test_data, test_labels, c="g", s=4, label="Testing data")

File ...:3684, in scatter(x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, edgecolors, plotnonfinite, data, **kwargs)
   3665 @_copy_docstring_and_deprecators(Axes.scatter)
   3666 def scatter(
   3667     x: float | ArrayLike,
   (...)
   3682     **kwargs,
   3683 ) -> PathCollection:
-> 3684     __ret = gca().scatter(
   3685         x,
   3686         y,
   3687         s=s,
   3688         c=c,
   3689         marker=marker,
...
   4654 if s is None:
   4655     s = (20 if mpl.rcParams['_internal.classic_mode'] else
   4656          mpl.rcParams['lines.markersize'] ** 2.0)

ValueError: x and y must be the same size

张量形状信息

X_tr.size() # torch.Size([275, 24])
X_t.size() # torch.Size([69, 24])
y_tr.size() # torch.Size([1, 275])
y_t.size() # torch.Size([69, 1])

问题原因与解决方案

核心问题

  1. 标签维度不匹配:y_tr是[1,275]的二维张量,train_data[:,0]是[275]的一维张量,取y_tr[:,0]得到仅1个元素的张量,与x轴数据长度不一致;y_t是[69,1],同样存在维度不匹配问题。
  2. 测试数据维度错误:test_data是[69,24]的二维张量,plt.scatter要求x轴输入一维数据,直接传入会导致维度兼容问题。

修正步骤

  1. 统一标签为一维张量:使用squeeze()或flatten()去除冗余维度。
  2. 测试数据取单一特征列:和训练数据保持一致,选择某一列(如第一列test_data[:,0])作为x轴数据。

修正后的完整代码

def plot_predictions(train_data=X_tr,
                     train_labels=y_tr,
                     test_data=X_t,
                     test_labels=y_t,
                     predictions=None):
  
    plt.figure(figsize=(10, 7))

    # 处理训练数据和标签维度
    train_x = train_data[:, 0]
    train_y = train_labels.squeeze()  # 转为一维 [275]

    plt.scatter(train_x, train_y, c="b", s=4, label="Training data")

    # 处理测试数据和标签维度
    test_x = test_data[:, 0]
    test_y = test_labels.squeeze()  # 转为一维 [69]

    plt.scatter(test_x, test_y, c="g", s=4, label="Testing data")

    if predictions is not None:
        pred_y = predictions.squeeze()  # 确保预测结果也是一维
        plt.scatter(test_x, pred_y, c="r", s=4, label="Predictions")
        
        correct = (pred_y == test_y).sum().item()
        total = test_y.size(0)
        accuracy = correct / total
        plt.title(f"Accuracy: {accuracy:.2f}")

    # Show the legend
    plt.legend(prop={"size": 14})
    plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 11:47:02