Python绘制多元线性回归模型时薪资预测值全为0的问题排查
问题
使用Python构建多元线性回归模型时出现异常:作为因变量的薪资(依赖年龄、工作年限等特征)预测值全部为0,但实际薪资范围应为30000至50000,与输出结果不符。以下是实现代码及异常可视化图:
# all required libraries import pandas as pd import warnings import numpy as np # For data visualizing import seaborn as sns #%matplotlib notebook import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D #%matplotlib inline %matplotlib widget # For building the required model from sklearn import linear_model df = pd.read_csv('ml_data_salary.csv') # Plotting a 3-D plot for visualizing the Multiple Linear Regression Model # Preparing the data X = df[['age', 'YearsExperience']].values.reshape(-1,2) Y = df['Salary'] # Create range for each dimension x = X[:, 0] y = X[:, 1] z = Y xx_pred = np.linspace(25, 40, 30) # range of age values yy_pred = np.linspace(1, 10, 30) # range of experience values xx_pred, yy_pred = np.meshgrid(xx_pred, yy_pred) model_viz = np.array([xx_pred.flatten(), yy_pred.flatten()]).T # Predict using model built on previous step ols = linear_model.LinearRegression() model1 = ols.fit(X, Y) predicted = model1.predict(model_viz) # Evaluate model by using it's R^2 score r2 = model.score(X, Y) # Plot model visualization plt.style.use('default') fig = plt.figure(figsize=(12, 4)) ax1 = fig.add_subplot(131, projection='3d') ax2 = fig.add_subplot(132, projection='3d') ax3 = fig.add_subplot(133, projection='3d') axes = [ax1, ax2, ax3] for ax in axes: ax.plot(x, y, z, color='k', zorder=15, linestyle='none', marker='o', alpha=0.5) ax.scatter(xx_pred.flatten(), yy_pred.flatten(), predicted, facecolor=(0,0,0,0), s=20, edgecolor='#70b3f0') ax.set_xlabel('Age', fontsize=12) ax.set_ylabel('Experience', fontsize=12) ax.set_zlabel('Salary', fontsize=12) ax.locator_params(nbins=4, axis='x') ax.locator_params(nbins=5, axis='x') ax1.view_init(elev=27, azim=112) ax2.view_init(elev=16, azim=-51) ax3.view_init(elev=60, azim=165) fig.suptitle('Multi-Linear Regression Model Visualization ($R^2 = %.2f$)' % r2, fontsize=15, color='k') fig.tight_layout()

问题排查与解决
核心问题分析
- 模型调用错误:计算R²得分时,代码使用了未定义的
model变量,实际训练完成的模型是model1,这会导致报错或异常结果(若环境中存在残留的同名变量)。 - 数据验证缺失:预测值全为0的最可能原因是训练数据异常——要么
Salary列被误读为全0,要么X特征与Y之间完全无线性相关性,或者数据读取时出现错误。 - 坐标轴设置重复:代码中重复对x轴设置刻度数量,其中一处应为y轴。
修正步骤
- 验证数据正确性:添加数据统计打印,确认
Salary列的数值范围是否符合预期(30000-50000)。 - 修正模型调用:将
model.score(X, Y)替换为model1.score(X, Y),确保使用训练好的模型计算R²得分。 - 检查模型参数:打印模型的系数和截距,确认模型是否正常训练(若系数和截距全为0,说明数据存在问题)。
- 修复坐标轴设置:将重复的x轴刻度设置改为y轴。
修正后代码
# 导入所需库 import pandas as pd import warnings import numpy as np # 可视化相关库 import seaborn as sns import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D %matplotlib widget # 模型构建库 from sklearn import linear_model df = pd.read_csv('ml_data_salary.csv') # 验证数据是否正确加载 print("数据基本统计信息:") print(df.describe()) # 准备训练数据 X = df[['age', 'YearsExperience']].values Y = df['Salary'] # 生成预测用的特征网格 x = X[:, 0] y = X[:, 1] z = Y xx_pred = np.linspace(25, 40, 30) # 年龄范围 yy_pred = np.linspace(1, 10, 30) # 工作年限范围 xx_pred, yy_pred = np.meshgrid(xx_pred, yy_pred) model_viz = np.array([xx_pred.flatten(), yy_pred.flatten()]).T # 训练模型并预测 ols = linear_model.LinearRegression() model1 = ols.fit(X, Y) # 打印模型参数,确认训练状态 print("\n模型系数:", model1.coef_) print("模型截距:", model1.intercept_) predicted = model1.predict(model_viz) # 计算模型R²得分 r2 = model1.score(X, Y) print("\n模型R²得分:", r2) # 绘制可视化图 plt.style.use('default') fig = plt.figure(figsize=(12, 4)) ax1 = fig.add_subplot(131, projection='3d') ax2 = fig.add_subplot(132, projection='3d') ax3 = fig.add_subplot(133, projection='3d') axes = [ax1, ax2, ax3] for ax in axes: ax.plot(x, y, z, color='k', zorder=15, linestyle='none', marker='o', alpha=0.5) ax.scatter(xx_pred.flatten(), yy_pred.flatten(), predicted, facecolor=(0,0,0,0), s=20, edgecolor='#70b3f0') ax.set_xlabel('Age', fontsize=12) ax.set_ylabel('Experience', fontsize=12) ax.set_zlabel('Salary', fontsize=12) ax.locator_params(nbins=4, axis='x') ax.locator_params(nbins=5, axis='y') # 修正重复设置x轴的问题 ax1.view_init(elev=27, azim=112) ax2.view_init(elev=16, azim=-51) ax3.view_init(elev=60, azim=165) fig.suptitle(f'Multi-Linear Regression Model Visualization ($R^2 = {r2:.2f}$)', fontsize=15, color='k') fig.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Muhammad Ali Siddiqui
相关产品推荐
相关产品推荐

