PyTorch多项式回归模型绘图报错:x与y维度不匹配
问题分析与解决方案
核心问题
- 绘图函数内部硬编码全局变量(
X_train、y_train),未使用定义的参数(train_features、train_labels),导致外部传参完全无效。 - 调用函数时传参错误:将
y_train的数组传给了train_features参数,造成x轴与y轴数据来源混乱。 - 未统一将二维张量(
[n,1])转换为一维数组,matplotlib对二维张量的维度处理易引发尺寸不匹配报错。
修正后的绘图函数
import matplotlib.pyplot as plt def plot_predictions(train_features=X_train, train_labels=y_train, test_features=X_test, test_labels=y_test, predictions=None): plt.figure(figsize=(10,7)) # 统一处理张量/数组:转numpy并压缩为一维 def process_data(data): if hasattr(data, 'detach'): # 判断是否为PyTorch张量 data = data.detach().numpy() return data.squeeze() # 移除维度为1的轴 # 处理所有输入数据 train_x = process_data(train_features) train_y = process_data(train_labels) test_x = process_data(test_features) test_y = process_data(test_labels) plt.scatter(train_x, train_y, c="g", label="training data") plt.scatter(test_x, test_y, c="b", label="testing data") if predictions is not None: pred = process_data(predictions) plt.scatter(test_x, pred, c="r", label="predictions") plt.legend(prop={"size":14})
正确调用方式
# 方式1:直接调用(函数内部自动处理张量) plot_predictions() # 方式2:手动传入预处理后的numpy数组(可选) plot_predictions(train_features=X_train.detach().numpy(), train_labels=y_train.detach().numpy(), test_features=X_test.detach().numpy(), test_labels=y_test.detach().numpy(), predictions=y_preds.detach().numpy())
关键修复说明
- 替换硬编码变量:将函数内的
X_train、y_train替换为定义的参数,确保外部传参能生效。 - 统一数据处理:新增
process_data函数,自动识别并转换PyTorch张量为numpy数组,同时压缩冗余维度,彻底解决尺寸不匹配问题。 - 修正传参逻辑:调用时严格对应参数含义,
train_features对应x轴数据,train_labels对应y轴数据,避免传参错位。
内容的提问来源于stack exchange,提问作者Bernardo Géo Cometto
相关产品推荐
相关产品推荐

