如何在scipy.optimize.curve_fit()中用kwargs传递非拟合参数
解决scipy.optimize.curve_fit传递额外参数的问题
问题场景
在衍射图案数据拟合场景中,已知缝数n,需要拟合曲线参数a和b,同时要将n传入拟合函数。尝试用**kwargs实现时触发报错,简化示例代码如下:
import numpy as np import scipy import matplotlib.pyplot as plt def curve(x,a,b,**kwargs): n = kwargs["n"] return a*np.sin(n*x)+b*np.cos(n*x) x = np.linspace(-5,5,1000) y = np.random.normal(loc=curve(x, 4, 3, n=2), scale=0.2, size=None) result = scipy.optimize.curve_fit(curve, x, y, n = 2) y2 = curve(x, *result[0], n=2) plt.plot(x, y2) plt.plot(x,y) plt.show()
运行后出现如下错误:
File "C:\Users\HP\OneDrive\Documents\Uni\lab year 2\diffraction\kwargs.py", line 13, in <module> result = scipy.optimize.curve_fit(curve, x, y, n = 2) File "C:\Users\HP\anaconda3\lib\site-packages\scipy\optimize\_minpack_py.py", line 834, in curve_fit res = leastsq(func, p0, Dfun=jac, full_output=1, **kwargs) TypeError: leastsq() got an unexpected keyword argument 'n'
错误原因
scipy.optimize.curve_fit中直接传入的**kwargs会被传递给底层的leastsq函数,而非自定义的curve函数。leastsq不识别n这个参数,因此抛出错误。
解决方法
方法1:使用curve_fit的args参数传递额外参数
curve_fit专门提供了args参数用于传递自定义拟合函数的额外参数,这是官方推荐的方式:
import numpy as np import scipy.optimize import matplotlib.pyplot as plt def curve(x, a, b, n): return a*np.sin(n*x)+b*np.cos(n*x) x = np.linspace(-5,5,1000) y = np.random.normal(loc=curve(x, 4, 3, 2), scale=0.2, size=None) # 将n以元组形式传入args result = scipy.optimize.curve_fit(curve, x, y, args=(2,)) y2 = curve(x, *result[0], 2) plt.plot(x, y2, label='拟合曲线') plt.plot(x, y, alpha=0.5, label='原始数据') plt.legend() plt.show()
方法2:用lambda函数包装原函数
如果坚持想用**kwargs,可以通过lambda函数封装,提前绑定n参数:
import numpy as np import scipy.optimize import matplotlib.pyplot as plt def curve(x,a,b,**kwargs): n = kwargs["n"] return a*np.sin(n*x)+b*np.cos(n*x) x = np.linspace(-5,5,1000) y = np.random.normal(loc=curve(x, 4, 3, n=2), scale=0.2, size=None) # 用lambda包装原函数,固定n参数 fit_func = lambda x, a, b: curve(x, a, b, n=2) result = scipy.optimize.curve_fit(fit_func, x, y) y2 = curve(x, *result[0], n=2) plt.plot(x, y2, label='拟合曲线') plt.plot(x, y, alpha=0.5, label='原始数据') plt.legend() plt.show()
方法3:给函数设置默认参数
如果拟合过程中n的值固定,直接给curve函数设置n的默认值即可:
import numpy as np import scipy.optimize import matplotlib.pyplot as plt def curve(x,a,b,n=2): return a*np.sin(n*x)+b*np.cos(n*x) x = np.linspace(-5,5,1000) y = np.random.normal(loc=curve(x, 4, 3), scale=0.2, size=None) result = scipy.optimize.curve_fit(curve, x, y) y2 = curve(x, *result[0]) plt.plot(x, y2, label='拟合曲线') plt.plot(x, y, alpha=0.5, label='原始数据') plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者Adam Eaton
相关产品推荐
相关产品推荐

