使用scipy curve_fit拟合返回DataFrame的函数时遇索引越界错误
自定义函数使用scipy.curve_fit拟合时触发"list index out of range"错误
问题描述
我用基于pandas实现的自定义函数拟合数据集,固定参数时函数能正常运行并绘图,但调用scipy的curve_fit优化参数时,出现“list index out of range”错误。
可正常运行的测试代码
import numpy as np import pandas as pd import matplotlib.pyplot as plt ydata=np.array([1.2,1.21,1.2,1.19,1.21,1.22,1.8,2.47,2.49,2.49,2.5]) xdata = np.linspace(0,1,num=11).round(2) def halfTrapz(x, m, a, tau1, tau2): y = np.zeros(len(x)) dfy = pd.DataFrame(list(zip(x,y)),columns=['x','y']) ta1=dfy.index[dfy['x']==tau1].tolist() ta2=dfy.index[dfy['x']==tau2].tolist() # ta1=list(np.array(ta1)+1) # ta2=list(np.array(ta2)+1) #In order to consider ta2 in [:ta2] dfy.iloc[:ta1[0],1] = a b = a - m*dfy.iloc[ta1[0],0] dfy.iloc[ta1[0]:ta2[0],1] = m * dfy.iloc[ta1[0]:ta2[0],0] + b dfy.iloc[ta2[0]:,1] = m * dfy.iloc[ta2[0],0] + b return dfy['y'] z=(halfTrapz(xdata, 5,1.2,0.5,0.7)) plt.plot(xdata,z,'g--') plt.plot(xdata,ydata)
出错的curve_fit调用代码
from scipy.optimize import curve_fit popt, pcov = curve_fit(halfTrapz, xdata, ydata) print(popt) print(pcov) plt.plot(xdata, func(xdata, *popt), 'r-')
错误原因
curve_fit在参数优化过程中会尝试大量参数组合,其中tau1或tau2的值大概率不会精确匹配xdata中的元素,导致dfy.index[dfy['x']==tau1].tolist()返回空列表。此时访问ta1[0]或ta2[0]就会触发索引越界错误。而手动测试时用的tau1=0.5、tau2=0.7刚好在xdata里,所以不会出错。
另外,浮点数值的精确匹配本身就存在风险,哪怕tau1理论上等于xdata中的某个值,也可能因浮点精度问题匹配失败。
解决方案
放弃用pandas做索引匹配,改用numpy的向量化操作,通过查找x中第一个大于等于tau1/tau2的位置来确定分段点,避免精确匹配的问题。同时这种方式效率更高,更适合优化过程中的大量调用。
修改后的函数示例:
def halfTrapz(x, m, a, tau1, tau2): # 确保tau1 <= tau2,避免分段顺序混乱 tau1, tau2 = sorted([tau1, tau2]) # 初始化y数组 y = np.full_like(x, a) # 找到第一个x >= tau1的索引 idx1 = np.argmax(x >= tau1) # 找到第一个x >= tau2的索引 idx2 = np.argmax(x >= tau2) # 计算第二段的线性表达式 b = a - m * x[idx1] y[idx1:idx2] = m * x[idx1:idx2] + b # 第三段保持第二段终点的值 y[idx2:] = m * x[idx2] + b return y
修改后再调用curve_fit即可正常运行,同时也避免了pandas操作带来的额外开销。
内容的提问来源于stack exchange,提问作者tulips
相关产品推荐
相关产品推荐

