You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 07:54:54