多分类模型预测结果绘图遇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])
问题原因与解决方案
核心问题
- 标签维度不匹配:
y_tr是[1,275]的二维张量,train_data[:,0]是[275]的一维张量,取y_tr[:,0]得到仅1个元素的张量,与x轴数据长度不一致;y_t是[69,1],同样存在维度不匹配问题。 - 测试数据维度错误:
test_data是[69,24]的二维张量,plt.scatter要求x轴输入一维数据,直接传入会导致维度兼容问题。
修正步骤
- 统一标签为一维张量:使用
squeeze()或flatten()去除冗余维度。 - 测试数据取单一特征列:和训练数据保持一致,选择某一列(如第一列
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
相关产品推荐
相关产品推荐

