高次多项式拟合中指定Y值对应的X值计算问题(Python)
高次多项式求根问题:
roots函数异常及优化方案 问题描述
尝试用Python的numpy库roots函数计算高次多项式在指定Y值(如0.8)对应的X值,流程为:用polyfit生成60次多项式系数,构造P(x)-y的多项式后调用roots求根,过滤虚数和超出范围的解。但结果仅一个X值正确,右侧出现大量错误解,即使降为30次问题依然存在。
复现代码:
import numpy as np from numpy.polynomial import Polynomial as poly import matplotlib.pyplot as plt def main(): # Declare sample data dataX = [0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, 6.5, 7.0, 7.5, 9.5, 10.0, 10.5, 11.0, 11.5, 12.0, 12.5, 13.0, 13.5, 14.0, 14.5, 15.0, 15.5, 16.0, 16.5, 17.0] dataY = [0.0, -0.008539747658600522, -0.06274870291613317, -0.19444406675443215, -0.36136059621877487, -0.600774409182093, -0.9162127771014803, -1.323039561600279, -1.8023953285501042, -2.3960659331052065, -3.1540213754489517, -4.067701143485799, -5.252273439144881, -6.771579806673188, -8.841718736169389, -11.702002554463427, -10.622244389289413, -7.093229874530799, -4.658886103196484, -2.6582045851707403, -1.046251560907664, 0.26783134185570256, 1.3575816391972544, 2.262728874235532, 2.990654481815292, 3.557426755904812, 3.955698729737411, 4.191874958827449, 4.253204560527751, 4.116558660927984, 3.7485922103648823, 3.1198702528381874] polyDegree = 60 # Perform polynomial fit (ascending coefficent order) polyCoeffs = np.flip(np.polyfit(dataX, dataY, polyDegree)) # Find all X-values at the desired Y-value of polynomial y = 0.8 x = (poly(polyCoeffs) - y).roots().tolist() # Filter out all found X-values that are imaginary or beyond desired limits (0, 17) i = 0 while (i < len(x)): if (x[i].imag != 0) or (x[i].real < 0) or (x[i].real > 17): del x[i] i -= 1 else: x[i] = x[i].real i += 1 # Plot data, polynomial and found X-values plt.xlim(-1, 18) plt.ylim(-20, 20) polyDataY = [] for i in range(len(dataX)): value = 0 for j in range(len(polyCoeffs)): value += polyCoeffs[j] * pow(dataX[i], j) polyDataY.append(value) plt.scatter(dataX, dataY, c = "dodgerblue", label = "Original Data") plt.plot(dataX, polyDataY, color = "orange", label = "Polynomial Fit") plt.axhline(y, color = "red", label = "Desired Y-value") for i in range(len(x)): plt.axvline(x[i], color = "forestgreen") plt.axvline(99999, color = "forestgreen", label = "Found X-values") plt.legend() plt.show() plt.close() plt.clf() if (__name__ == '__main__'): main()
问题原因
- 高次多项式的数值不稳定性:60次属于超高次多项式,
polyfit计算系数时会因高次幂的放大效应产生严重数值误差,微小的系数偏差会导致多项式值在样本点外剧烈波动,求根结果自然偏离真实解。 roots函数的局限性:数值求根算法对高次多项式的系数误差极为敏感,当多项式系数量级差异大或存在接近的根时,容易计算出大量虚假的实根。- 严重过拟合:样本仅32个点,用60次多项式拟合会完全拟合噪声,拟合出的多项式在样本区间内波动极大,根本无法反映数据的真实趋势,求根必然出现大量错误解。
优化方案
1. 选择合适的多项式次数
避免使用过高次数,可通过交叉验证确定最优次数(比如尝试5-10次),平衡拟合效果和数值稳定性。示例修改:
polyDegree = 8 # 改用低次多项式
2. 用非线性优化代替多项式求根
无需构造新多项式,直接对P(x) - y = 0在指定区间内用单变量求根函数求解,稳定性更强。使用scipy.optimize.root_scalar示例:
import numpy as np from numpy.polynomial import Polynomial as poly import matplotlib.pyplot as plt from scipy.optimize import root_scalar def main(): dataX = [0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, 6.5, 7.0, 7.5, 9.5, 10.0, 10.5, 11.0, 11.5, 12.0, 12.5, 13.0, 13.5, 14.0, 14.5, 15.0, 15.5, 16.0, 16.5, 17.0] dataY = [0.0, -0.008539747658600522, -0.06274870291613317, -0.19444406675443215, -0.36136059621877487, -0.600774409182093, -0.9162127771014803, -1.323039561600279, -1.8023953285501042, -2.3960659331052065, -3.1540213754489517, -4.067701143485799, -5.252273439144881, -6.771579806673188, -8.841718736169389, -11.702002554463427, -10.622244389289413, -7.093229874530799, -4.658886103196484, -2.6582045851707403, -1.046251560907664, 0.26783134185570256, 1.3575816391972544, 2.262728874235532, 2.990654481815292, 3.557426755904812, 3.955698729737411, 4.191874958827449, 4.253204560527751, 4.116558660927984, 3.7485922103648823, 3.1198702528381874] polyDegree = 8 polyCoeffs = np.flip(np.polyfit(dataX, dataY, polyDegree)) p = poly(polyCoeffs) y_target = 0.8 # 定义目标函数 def func(x): return p(x) - y_target # 在可能的区间内找根,根据数据趋势划分区间 intervals = [(11, 12), (16, 17)] roots = [] for a, b in intervals: try: res = root_scalar(func, bracket=[a, b], method='brentq') if res.converged: roots.append(res.root) except ValueError: continue # 区间内无实根则跳过 # 绘图部分 plt.xlim(-1, 18) plt.ylim(-20, 20) polyDataX = np.linspace(0, 17, 1000) polyDataY = p(polyDataX) plt.scatter(dataX, dataY, c="dodgerblue", label="Original Data") plt.plot(polyDataX, polyDataY, color="orange", label="Polynomial Fit") plt.axhline(y_target, color="red", label="Desired Y-value") for root in roots: plt.axvline(root, color="forestgreen") plt.axvline(99999, color="forestgreen", label="Found X-values") plt.legend() plt.show() if __name__ == '__main__': main()
3. 改用样条插值
样条插值分段拟合,数值稳定性远高于高次多项式,适合非线性数据。使用scipy.interpolate.UnivariateSpline示例:
import numpy as np import matplotlib.pyplot as plt from scipy.interpolate import UnivariateSpline def main(): dataX = [0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, 6.5, 7.0, 7.5, 9.5, 10.0, 10.5, 11.0, 11.5, 12.0, 12.5, 13.0, 13.5, 14.0, 14.5, 15.0, 15.5, 16.0, 16.5, 17.0] dataY = [0.0, -0.008539747658600522, -0.06274870291613317, -0.19444406675443215, -0.36136059621877487, -0.600774409182093, -0.9162127771014803, -1.323039561600279, -1.8023953285501042, -2.3960659331052065, -3.1540213754489517, -4.067701143485799, -5.252273439144881, -6.771579806673188, -8.841718736169389, -11.702002554463427, -10.622244389289413, -7.093229874530799, -4.658886103196484, -2.6582045851707403, -1.046251560907664, 0.26783134185570256, 1.3575816391972544, 2.262728874235532, 2.990654481815292, 3.557426755904812, 3.955698729737411, 4.191874958827449, 4.253204560527751, 4.116558660927984, 3.7485922103648823, 3.1198702528381874] y_target = 0.8 # 构建样条插值,s参数控制平滑度,0为完全拟合 spl = UnivariateSpline(dataX, dataY, s=0.1) # 找根 roots = spl.roots(y_target) # 过滤范围外的根 roots = [r for r in roots if 0 <= r <=17] # 绘图 plt.xlim(-1, 18) plt.ylim(-20, 20) x_plot = np.linspace(0,17,1000) plt.scatter(dataX, dataY, c="dodgerblue", label="Original Data") plt.plot(x_plot, spl(x_plot), color="orange", label="Spline Fit") plt.axhline(y_target, color="red", label="Desired Y-value") for root in roots: plt.axvline(root, color="forestgreen") plt.axvline(99999, color="forestgreen", label="Found X-values") plt.legend() plt.show() if __name__ == '__main__': main()
内容的提问来源于stack exchange,提问作者Runsva
相关产品推荐
相关产品推荐

