如何泛化拟合函数,使SciPy curve_fit自动推断输入参数数量
问题描述
我开发了一个程序,可生成带有可变数量未知变量的SymPy lambdify符号函数。我希望无需显式传递变量的准确数量,即可使用curve_fit工具对该函数进行拟合。
现有可运行代码如下:
import numpy as np import matplotlib as mpl mpl.use('TkAgg') import matplotlib.pyplot as plt from scipy.optimize import curve_fit from scipy import integrate import sympy as sy def step_function(t, start, end): return 1 * ((t >= start) & (t < end)) def gen_data(x_data, model, gradients, offset=0, noise=0.1): # generate some fake data y = [model(x, offset, gradients[0], gradients[1], gradients[2]) for x in x_data] y = np.asarray(y, dtype=np.float32) if noise: y += np.random.normal(0, 1, 100) return y t, a = sy.symbols('t, a') grad_symbols = sy.symbols('b, c, d') x = np.arange(0, 100) gradient_changes = [0, 30, 60, 100] gradients = [1, 2, -4] grad = [] for t1, t2, LFD in zip(iter(gradient_changes[::1]), iter(gradient_changes[1::1]), iter(grad_symbols)): grad.append(((t - t1) * LFD) * sy.SingularityFunction(t, t1, gradient_changes[-1])) data_model = sy.lambdify((t, a) + grad_symbols, a + sum(grad), ({'SingularityFunction': step_function}, 'numpy')) y = gen_data(x, data_model, gradients, noise=0.1) fig, ax = plt.subplots() # Create a figure containing a single axes. ax.plot(x, y, label='Data') popt, pcov = curve_fit(data_model, x, y) fit_line = [data_model(time, popt[0], popt[1], popt[2], popt[3]) for time in x] ax.plot(x, fit_line, label='Fit')
当前curve_fit可识别lambdify表达式的参数数量,但由于数据非瞬时性,我实际需要拟合函数在每个点附近的积分值。我尝试使用如下带星号表达式的拟合函数:
def fit_func(x, *inputs): # find the average value of the model about each point return [integrate.quad(data_model, _x - 1 / 2, _x + 1 / 2, args=inputs)[0] for _x in x]
并通过popt, pcov = curve_fit(fit_func, x, y)进行拟合,但目前报错:ValueError: Unable to determine number of fit parameters。请问该方法是否可行?如何解决该问题?
解决方案
这个方法是可行的,报错的核心原因是curve_fit无法通过带*inputs的函数自动推断需要拟合的参数数量。解决思路是给curve_fit提供初始参数猜测值p0,让它明确知道参数的数量,具体实现如下:
- 动态获取参数数量:从
data_model的定义可知,参数包含偏移量a加上所有梯度符号,可通过1 + len(grad_symbols)动态计算,避免硬编码。 - 生成初始参数猜测:用全1数组或基于数据的合理值作为初始猜测(比如根据原始数据的偏移和梯度范围调整)。
- 调用curve_fit时传入p0:
修改后的拟合代码如下:
# 动态计算需要拟合的参数总数 param_count = 1 + len(grad_symbols) # 生成初始参数猜测值 p0 = np.ones(param_count) # 传入p0调用curve_fit popt, pcov = curve_fit(fit_func, x, y, p0=p0) # 绘制积分后的拟合曲线 fit_line = fit_func(x, *popt) ax.plot(x, fit_line, label='Integrated Fit') ax.legend() plt.show()
额外优化建议:如果数据量较大,循环调用integrate.quad会影响效率,可考虑将积分逻辑向量化,或者预先推导积分后的解析表达式(利用SymPy的积分功能直接生成积分后的符号函数,再转为lambdify函数),进一步提升拟合速度。
内容的提问来源于stack exchange,提问作者Kurtresponse
相关产品推荐
相关产品推荐

