未用numpy实现线性回归,axline绘制回归线出现异常的解决方法
手动实现线性回归后无法正确绘制回归线
我出于学习目的手动实现简单线性回归,计算出斜率b和截距a,代码能正常绘制散点图与预测点,但添加plt.axline(xy1=(0, a), slope=b, linestyle="--", color="k")绘制回归线时出现异常结果,请问如何修正?
原始代码:
import matplotlib.pyplot as plt x = [1.5, 1.6, 1.7, 1.8, 1.9, 2.0, 2.1, 2.2] y = [60, 62, 64, 66, 68, 70, 72, 74] n = len(x) sx = sum(x) sy = sum(y) sxy = sum([x[i] * y[i] for i in range(n)]) sx2 = sum([x[i] ** 2 for i in range(n)]) b = (n * sxy - sx * sy) / (n * sx2 - sx ** 2) a = (sy / n) - b * (sx / n) def predict_peso(altura): return a + b * altura altura_prev = 1.75 peso_prev = predict_peso(altura_prev) plt.plot(altura_prev, peso_prev, marker="o", markeredgecolor="red", markerfacecolor="green") plt.scatter(x, y) plt.show()
问题原因
你计算出的回归线公式为 y = 30 + 20x,(0, a) 即(0,30)这个点远低于当前散点的y值区间(60-74)。plt.axline 默认会将线条延伸至整个坐标轴范围,这会导致图表的y轴被大幅拉伸,使得回归线看起来位置异常,同时散点的显示被压缩。
修正方案
不需要从(0,a)绘制回归线,而是使用数据x范围内的两个端点来生成线条,这样线条会贴合散点的显示范围,不会干扰坐标轴的缩放。
修正后的完整代码:
import matplotlib.pyplot as plt x = [1.5, 1.6, 1.7, 1.8, 1.9, 2.0, 2.1, 2.2] y = [60, 62, 64, 66, 68, 70, 72, 74] n = len(x) sx = sum(x) sy = sum(y) sxy = sum([x[i] * y[i] for i in range(n)]) sx2 = sum([x[i] ** 2 for i in range(n)]) b = (n * sxy - sx * sy) / (n * sx2 - sx ** 2) a = (sy / n) - b * (sx / n) def predict_peso(altura): return a + b * altura altura_prev = 1.75 peso_prev = predict_peso(altura_prev) plt.plot(altura_prev, peso_prev, marker="o", markeredgecolor="red", markerfacecolor="green") plt.scatter(x, y) # 修正:使用数据x的最小和最大值计算对应的y值,绘制回归线 x_min = min(x) x_max = max(x) y_min = predict_peso(x_min) y_max = predict_peso(x_max) plt.plot([x_min, x_max], [y_min, y_max], linestyle="--", color="k") plt.show()
补充说明
如果你坚持使用plt.axline,可以通过设置坐标轴范围来限制显示,比如添加plt.ylim(min(y)-2, max(y)+2),但用数据端点绘制线条的方式更直观且适配性更好。
内容的提问来源于stack exchange,提问作者celsowm
相关产品推荐
相关产品推荐

