如何修改现有Python代码生成符合标准多项式回归公式的数据集
多项式数据集生成代码调整方案
你当前的代码和目标公式主要有四处差异,对应调整如下:
- 修正随机误差项的均值:目标公式中e是均值为0的随机误差,你现有代码里误差项为
np.random.normal(-3, 3, n)均值为-3,需要调整为np.random.normal(0, 标准差, n),标准差可按需设置,比如保留原有的3就写np.random.normal(0, 3, n)。 - 显式定义多项式系数数组:对应公式里的B₀(截距项)、B₁、B₂...Bₙ,你可以把系数统一放在一个列表里,方便灵活调整多项式阶数,比如3阶多项式可以定义为
beta = [B0, B1, B2, B3],你原有逻辑里的系数对应为beta = [0, 1, -2, 0.5],如果需要加截距项直接修改beta第一个元素即可。 - 优化y的生成逻辑适配任意阶数:不需要手动逐阶写幂次计算,可通过循环自动累加各阶项,适配不同阶数的多项式生成需求。
- 删除冗余变量:现有代码中定义的
m和b没有被使用,如果没有后续用途可以直接删除,若需要作为系数使用可对应合并到beta数组中。
调整后的完整示例代码如下:
import numpy as np import matplotlib.pyplot as plt def generate_poly_dataset(beta, n=500, error_std=3): # beta为多项式系数数组,顺序为[B0, B1, B2, ..., Bk],对应k阶多项式 X = 2 - 3 * np.random.normal(0, 1, n) y = np.zeros(n) # 累加各阶多项式项 for order, b in enumerate(beta): y += b * (X ** order) # 添加均值为0的随机误差 y += np.random.normal(0, error_std, n) plt.scatter(X, y, s=10) plt.show() return X, y # 调用示例:生成和你原有逻辑一致的3阶多项式,仅误差均值改为0 # 对应系数B0=0, B1=1, B2=-2, B3=0.5 beta = [0, 1, -2, 0.5] X, y = generate_poly_dataset(beta)
内容的提问来源于stack exchange,提问作者Opps_0
相关产品推荐
相关产品推荐

