使用numpy.polynomial.Polynomial.fit拟合信号后绘图异常,求排查
多项式拟合曲线不符合预期的问题排查与解决
问题描述
使用numpy.polynomial.polynomial.Polynomial.fit函数拟合正弦信号,提取系数后自行代入多项式方程计算y值并绘图,得到的橙色拟合曲线严重偏离原始正弦信号(蓝色曲线),不符合预期。
问题代码
import math def getYValueFromCoeff(f,coeff_list): # low to high order y_plot_values=[] for j in range(len(f)): item_list= [] for i in range(len(coeff_list)): item= (coeff_list[i])*((f[j])**i) item_list.append(item) y_plot_values.append(sum(item_list)) print(len(y_plot_values)) return y_plot_values from numpy.polynomial import Polynomial as poly import numpy as np import matplotlib.pyplot as plt no_of_coef= 10 #original signal x = np.linspace(0, 0.01, 10) period = 0.01 y = np.sin(np.pi * x / period) #poly fit test1= poly.fit(x,y,no_of_coef) coeffs= test1.coef #print(test1.coef) coef_y= getYValueFromCoeff(x, test1.coef) #print(coef_y) plt.plot(x,y) plt.plot(x, coef_y)
错误原因
poly.fit函数默认会对输入的自变量x进行标准化处理(将其线性映射到[-1, 1]区间),目的是提升数值稳定性。因此返回的test1.coef是基于标准化后x的多项式系数,而非原始x。直接用原始x代入多项式计算,自然会得到偏离预期的结果。
解决方法
方法1:直接使用拟合对象计算y值(推荐)
不需要自行编写计算函数,拟合得到的test1本身就是可调用的多项式对象,直接传入原始x即可得到正确的拟合值:
import numpy as np from numpy.polynomial import Polynomial as poly import matplotlib.pyplot as plt no_of_coef= 10 # 原始信号 x = np.linspace(0, 0.01, 10) period = 0.01 y = np.sin(np.pi * x / period) # 多项式拟合 test1 = poly.fit(x, y, no_of_coef) # 直接用拟合对象计算y值 coef_y = test1(x) plt.plot(x, y, label='原始信号') plt.plot(x, coef_y, label='拟合曲线') plt.legend() plt.show()
方法2:关闭标准化,使用原始x的系数
在poly.fit时指定domain参数为原始x的取值范围,这样函数会跳过标准化步骤,返回基于原始x的系数,自定义计算函数就能正常工作:
import math import numpy as np from numpy.polynomial import Polynomial as poly import matplotlib.pyplot as plt def getYValueFromCoeff(f,coeff_list): # low to high order y_plot_values=[] for j in range(len(f)): item_list= [] for i in range(len(coeff_list)): item= coeff_list[i] * (f[j]**i) item_list.append(item) y_plot_values.append(sum(item_list)) return y_plot_values no_of_coef= 10 # 原始信号 x = np.linspace(0, 0.01, 10) period = 0.01 y = np.sin(np.pi * x / period) # 拟合时指定domain为原始x的范围,关闭标准化 test1 = poly.fit(x, y, no_of_coef, domain=[x.min(), x.max()]) coef_y = getYValueFromCoeff(x, test1.coef) plt.plot(x, y, label='原始信号') plt.plot(x, coef_y, label='拟合曲线') plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者MRR
相关产品推荐
相关产品推荐

