如何获取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.
相关产品推荐
相关产品推荐

