如何用Scipy.optimize.leastsq拟合倒数幂函数并规避异常?
解决scipy.optimize拟合倒数幂函数的数值问题
问题背景
需要用Python 3.8的scipy.optimize将XY数据拟合至函数y = -d/(x+s)^e,原代码使用leastsq但出现数值错误,拟合失败,收到除零、无效幂运算及达到最大迭代次数的警告,拟合曲线完全不符合预期。
原代码:
import numpy as np from scipy.optimize import leastsq import matplotlib.pyplot as plt def main(): # Declare example data dataX = [0, 0.5, 1, 1.5, 2, 2.5, 3, 3.5, 4, 4.5, 5, 5.5, 6, 6.5, 7, 7.5, 8, 8.5, 9, 9.5, 10, 10.5, 11, 11.5, 12, 12.5, 13, 13.5, 14.5, 15, 15.5, 16, 16.5, 17] dataY = [0.00000000e+00, 2.47076895e-03, 9.66638150e-03, 1.97670203e-02, 3.42218835e-02, 4.97943540e-02, 6.54004261e-02, 7.93484613e-02, 8.61083796e-02, 8.49273950e-02, 6.70327841e-02, 3.45946900e-02, -2.24559129e-02, -1.17023372e-01, -2.49244700e-01, -4.27837601e-01, -6.61529080e-01, -9.62240010e-01, -1.35215046e+00, -1.84177480e+00, -2.46428273e+00, -3.25198094e+00, -4.24238794e+00, -5.53021367e+00, -7.23278230e+00, -9.57603920e+00, -1.31432215e+01, -1.99465466e+01, -1.45862198e+01, -1.01563667e+01, -7.19977087e+00, -5.03349403e+00, -3.37202972e+00, -2.07966002e+00] # Fit to reciprocal exponential function guessDenom = 1 guessPow = 2 guessShift = -10 optimalFx = lambda args: -(args[0] / pow(dataX + args[2], args[1])) - dataY finalDenom, finalPow, finalShift = leastsq(optimalFx, [guessDenom, guessPow, guessShift])[0] dataYFitted = [] for x in dataX: dataYFitted.append(-finalDenom / pow(x + finalShift, finalPow)) # Plot original & fitted data plt.xlim([0, 17]) plt.ylim([-30, 30]) plt.plot(dataX, dataYFitted, label = "Fitted Function", color = "dodgerblue") plt.plot(dataX, dataY, label = "Original Data", color = "gray", alpha = 0.5) plt.legend() plt.show() plt.close() plt.clf() if (__name__ == "__main__"): main()
收到的警告:
RuntimeWarning: divide by zero encountered in divide optimalFx = lambda args: -(args[0] / pow(dataX + args[2], args[1])) - dataY RuntimeWarning: invalid value encountered in power optimalFx = lambda args: -(args[0] / pow(dataX + args[2], args[1])) - dataY RuntimeWarning: Number of calls to function has reached maxfev = 800. warnings.warn(errors[info][0], RuntimeWarning)
问题根源
- 数据类型错误:
dataX是Python列表,与数值args[2]相加时执行列表拼接而非逐元素数值运算,导致计算完全错误。 - 无效数值区间:拟合过程中参数
s可能让x+s趋近于0(引发除零)或变为负数(当e为非整数时,负数开幂得到NaN),导致数值崩溃。 - leastsq的局限性:
leastsq不支持直接设置参数搜索范围,无法约束参数进入有效区间。
解决方案
1. 改用支持参数约束的curve_fit
scipy.optimize.curve_fit可以通过bounds参数限制参数范围,接口更直观,适合曲线拟合场景。
2. 修正数据类型与拟合函数
- 将
dataX和dataY转为numpy数组,确保逐元素运算正确。 - 定义拟合函数时加入数值保护,避免
x+s进入无效区间。
3. 优化初始猜测值
根据数据趋势调整初始猜测:观察数据在x≈13后急剧下降,说明x+s在此处趋近于0,因此s的初始猜测应设为-13左右而非-10。
修改后的完整代码
import numpy as np from scipy.optimize import curve_fit import matplotlib.pyplot as plt def fit_func(x, d, e, s): # 数值保护:避免分母过小或为负,引发除零/无效幂运算 denominator = x + s denominator = np.where(denominator <= 1e-6, 1e-6, denominator) return -d / (denominator ** e) def main(): # 转为numpy数组,确保逐元素运算正确 dataX = np.array([0, 0.5, 1, 1.5, 2, 2.5, 3, 3.5, 4, 4.5, 5, 5.5, 6, 6.5, 7, 7.5, 8, 8.5, 9, 9.5, 10, 10.5, 11, 11.5, 12, 12.5, 13, 13.5, 14.5, 15, 15.5, 16, 16.5, 17]) dataY = np.array([0.00000000e+00, 2.47076895e-03, 9.66638150e-03, 1.97670203e-02, 3.42218835e-02, 4.97943540e-02, 6.54004261e-02, 7.93484613e-02, 8.61083796e-02, 8.49273950e-02, 6.70327841e-02, 3.45946900e-02, -2.24559129e-02, -1.17023372e-01, -2.49244700e-01, -4.27837601e-01, -6.61529080e-01, -9.62240010e-01, -1.35215046e+00, -1.84177480e+00, -2.46428273e+00, -3.25198094e+00, -4.24238794e+00, -5.53021367e+00, -7.23278230e+00, -9.57603920e+00, -1.31432215e+01, -1.99465466e+01, -1.45862198e+01, -1.01563667e+01, -7.19977087e+00, -5.03349403e+00, -3.37202972e+00, -2.07966002e+00]) # 优化初始猜测值,贴合数据趋势 initial_guess = [100, 2, -13] # 设置参数范围:d、e为正数,s确保x+s在所有数据点上有效 bounds = ([1e-3, 1e-3, -17], [1e5, 10, -12]) # 执行拟合 params, _ = curve_fit(fit_func, dataX, dataY, p0=initial_guess, bounds=bounds) final_d, final_e, final_s = params # 生成拟合数据 dataYFitted = fit_func(dataX, final_d, final_e, final_s) # 绘图展示 plt.xlim([0, 17]) plt.ylim([-30, 30]) plt.plot(dataX, dataYFitted, label="Fitted Function", color="dodgerblue") plt.plot(dataX, dataY, label="Original Data", color="gray", alpha=0.5) plt.legend() plt.show() plt.close() if __name__ == "__main__": main()
关键改动说明
- 数值保护:用
np.where将过小的x+s替换为1e-6,避免除零和无效幂运算。 - 参数约束:
bounds限制d和e为正数(符合函数物理意义),s的范围确保所有数据点的x+s处于有效区间。 - 初始猜测优化:根据数据趋势调整
s的初始值,让优化更快收敛到正确区间。
内容的提问来源于stack exchange,提问作者Runsva
相关产品推荐
相关产品推荐

