You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何获取sklearn PolynomialFeatures生成的多项式回归模型方程?

获取Sklearn多项式回归模型的方程

我需要从用sklearn.preprocessing.PolynomialFeatures生成的多项式回归模型中提取出方程,方便在其他程序里用这个模型做预测。多项式回归方程的形式类似:y = b₀ + b₁x + b₂x²。以下是我的实验代码:

import matplotlib.pyplot as plt
from sklearn.linear_model import LinearRegression
from sklearn.preprocessing import PolynomialFeatures


def viz_linear():
    plt.scatter(X, y, color='red')
    plt.plot(X, lin_reg.predict(X), color='blue')
    plt.show()
    return viz_linear()


def viz_polynomial():
    plt.scatter(X, y, color='red')
    plt.plot(X, pol_reg.predict(poly_reg.fit_transform(X)), color='blue')
    plt.show()
    return viz_polynomial()


if __name__ == '__main__':
    # 实验数据
    y = [0.0, 0.1, 0.2, 0.29, 0.39, 0.5, 0.62, 0.72, 0.82, 0.88, 1.05, 1.31, 1.61, 2.1, 2.48, 2.58, 2.88, 3.32, 3.66, 4.28, 4.54, 4.95, 5.09, 5.67, 7.09, 8.48, 9.03, 9.22, 10.0]
    X = [[0.0], [138.0], [259.0], [360.0], [439.0], [515.0], [558.0], [583.0], [606.0], [617.0], [648.0], [684.0], [715.0], [754.0], [776.0], [781.0], [795.0], [812.0], [824.0], [841.0], [847.0], [855.0], [858.0], [868.0], [888.0], [902.0], [906.0], [908.0], [914.0]]

    lin_reg = LinearRegression()
    lin_reg.fit(X, y)
    
    poly_reg = PolynomialFeatures(degree=10)
    X_poly = poly_reg.fit_transform(X)
    pol_reg = LinearRegression()
    pol_reg.fit(X_poly, y)
    
    viz_polynomial()

    print(pol_reg.predict(poly_reg.fit_transform([[914]])))

实现步骤

要提取方程,核心是获取模型的截距项和各项系数,再与多项式特征项对应:

  • 截距项b₀:对应pol_reg.intercept_
  • 各项系数b₁, b₂,...:对应pol_reg.coef_,注意跳过第0个元素(对应常数项,已单独放在截距中)
  • 多项式特征项:用poly_reg.get_feature_names_out()获取每个系数对应的表达式(如x, x^2)

代码实现

在你的代码末尾添加以下逻辑,即可输出完整多项式方程:

# 获取特征名称与模型系数
feature_names = poly_reg.get_feature_names_out()
coefficients = pol_reg.coef_
intercept = pol_reg.intercept_

# 构建方程片段
equation_parts = [f"{intercept:.6f}"]
for name, coef in zip(feature_names[1:], coefficients[1:]):
    # 处理正负号优化方程可读性
    if coef >= 0:
        equation_parts.append(f"+ {coef:.6f}*{name}")
    else:
        equation_parts.append(f"- {abs(coef):.6f}*{name}")

# 拼接并打印方程
polynomial_equation = "y = " + " ".join(equation_parts)
print("\n多项式回归方程:")
print(polynomial_equation)

说明

  • :.6f用于保留6位小数,可根据需求调整精度
  • 跳过feature_names[0]和coefficients[0]是因为第0个特征为1,对应截距项,已单独处理
  • 输出的方程可直接复制到其他程序中,用于手动计算预测值

内容的提问来源于stack exchange,提问作者Zach M.

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.15 21:45:17