同一画布双Linear Regression预测曲线显示异常问题咨询
线性回归预测曲线显示异常问题解决
问题描述
我想在同一画布的两个子图中,分别为两组数据构建并绘制不同的线性回归模型,但遇到了y1_pred的显示问题——它无法覆盖散点所在的整个y轴范围。
模型训练代码
model1 = LinearRegression().fit(x, y) intercept1 = model1.intercept_ slope1 = model1.coef_ y1_pred = intercept1 + slope1 * x model2 = LinearRegression().fit(x2, y2) intercept2 = model2.intercept_ slope2 = model2.coef_ y2_pred = intercept2 + slope2 * x2
绘图代码
axs[0].scatter(x, y, marker='o', color='aqua', label='1980-2000', edgecolors = 'black', linewidths = 1, s=500, zorder=4) axs[1].scatter(x2, y2, marker='o', color='magenta', label='2000-2024', edgecolors = 'black', linewidths = 1, s=500, zorder=4) axs[0].plot(x2, y2_pred, color='red', linewidth=4) axs[1].plot(x, y1_pred, color='red', linewidth=4)
异常显示效果

问题原因
你的绘图代码存在数据不匹配的核心问题:
- 第一个子图(
axs[0])对应的数据是x和y,但你用了另一组数据的x2来绘制y2_pred,导致回归线的x范围和散点的x范围完全不匹配,自然无法覆盖散点的y轴范围 - 第二个子图(
axs[1])同理,用x来绘制y1_pred,但子图内的散点是x2和y2,x轴范围不对应,直接导致回归线显示异常
解决方案
方案1:匹配对应数据绘制回归线
修正绘图代码,让每个子图的回归线使用对应模型训练时的x数据:
axs[0].scatter(x, y, marker='o', color='aqua', label='1980-2000', edgecolors = 'black', linewidths = 1, s=500, zorder=4) axs[1].scatter(x2, y2, marker='o', color='magenta', label='2000-2024', edgecolors = 'black', linewidths = 1, s=500, zorder=4) # 子图0绘制model1的回归线,使用训练数据x axs[0].plot(x, y1_pred, color='red', linewidth=4) # 子图1绘制model2的回归线,使用训练数据x2 axs[1].plot(x2, y2_pred, color='red', linewidth=4)
方案2:让回归线覆盖子图全x轴范围
如果希望回归线能延伸到子图的整个x轴范围(而非仅数据点的x范围),可以生成覆盖子图x轴全范围的序列来计算预测值:
import numpy as np axs[0].scatter(x, y, marker='o', color='aqua', label='1980-2000', edgecolors = 'black', linewidths = 1, s=500, zorder=4) axs[1].scatter(x2, y2, marker='o', color='magenta', label='2000-2024', edgecolors = 'black', linewidths = 1, s=500, zorder=4) # 对子图0生成覆盖x轴全范围的x序列 x_full0 = np.linspace(axs[0].get_xlim()[0], axs[0].get_xlim()[1], 100) y1_pred_full = intercept1 + slope1 * x_full0 axs[0].plot(x_full0, y1_pred_full, color='red', linewidth=4) # 对子图1生成覆盖x轴全范围的x序列 x_full1 = np.linspace(axs[1].get_xlim()[0], axs[1].get_xlim()[1], 100) y2_pred_full = intercept2 + slope2 * x_full1 axs[1].plot(x_full1, y2_pred_full, color='red', linewidth=4)
内容的提问来源于stack exchange,提问作者Gonzalo Martinez
相关产品推荐
相关产品推荐

