使用scipy curve_fit的absolute_sigma参数时触发ValueError错误
问题:curve_fit添加absolute_sigma参数触发ValueError错误
我编写了如下Python代码,用于通过模拟数据做线性拟合:
import numpy as np from scipy.optimize import curve_fit def lin_func(x, a, b): return a * x + b a = 0.005/4 b = 10000/4 N = 10000 position_list = np.array([-100,-50,-20,0,20,50,100]) def get_params(N): counts_list = np.zeros(len(position_list)) for i in range(len(position_list)): position_0 = position_list[i] counts = 0 for j in range(N): position = np.random.normal(position_0,10) p = 0.5+a+b*position*10**(-6) counts = counts + np.sum(np.random.binomial(1000, p, 1)) counts_list[i] = counts print(np.sqrt(counts_list)) popt, pcov = curve_fit(lin_func, position_list*10**(-6), counts_list, absolute_sigma = np.sqrt(counts_list)) return popt[1]/(N*1000) K = 100 params_list = np.zeros(K) for i in range(K): params_list[i] = get_params(N)-0.5 print(params_list[i])
程序运行后触发如下错误:
ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
该错误仅在添加absolute_sigma = np.sqrt(counts_list)参数时出现,移除该参数则运行正常,出错前打印的counts_list数值无异常。完整报错栈信息如下:
ValueError Traceback (most recent call last) Cell In[10], line 4 2 params_list = np.zeros(K) 3 for i in range(K): ----> 4 params_list[i] = get_params(N)-0.5 5 print(params_list[i]) Cell In[9], line 13, in get_params(N) 11 counts_list[i] = counts 12 print(np.sqrt(counts_list)) ---> 13 popt, pcov = curve_fit(lin_func, position_list*10**(-6), counts_list, absolute_sigma = np.sqrt(counts_list)) 14 return popt[1]/(N*1000) File ~/miniforge3/lib/python3.10/site-packages/scipy/optimize/_minpack_py.py:897, in curve_fit(f, xdata, ydata, p0, sigma, absolute_sigma, check_finite, bounds, method, jac, full_output, **kwargs) 895 pcov.fill(inf) 896 warn_cov = True --> 897 elif not absolute_sigma: 898 if ysize > p0.size: 899 s_sq = cost / (ysize - p0.size) ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
解决方案
问题出在参数传错了:
absolute_sigma是一个布尔值参数,作用是指定sigma参数传入的误差是绝对误差还是相对权重;- 你要传入的误差数组
np.sqrt(counts_list)应该传给sigma参数,而不是absolute_sigma。
修正后的curve_fit调用代码如下:
popt, pcov = curve_fit(lin_func, position_list*10**(-6), counts_list, sigma=np.sqrt(counts_list), absolute_sigma=True)
这样既传入了每个数据点的误差数组,又通过absolute_sigma=True声明这是绝对误差,就不会触发布尔值判断的歧义错误了。
内容的提问来源于stack exchange,提问作者Alex Marshall
相关产品推荐
相关产品推荐

