使用SciPy curve_fit拟合曲线:拐点处拟合不佳问题排查
SciPy curve_fit拟合带拐点曲线时拐点区域效果不佳的问题
我尝试用SciPy的curve_fit拟合如下方程:
$$y = 8A(0.1875*(x+B))^{1/3}$$
参数约束为:$A \in [0,1]$,$B \in [1,25]$。但拟合结果在曲线拐点处表现很差,以下是我的实现代码:
import matplotlib.pyplot as plt import numpy as np from scipy.optimize import curve_fit plt.close("all") def find_nearest(array,value): idx = (np.abs(array-value)).argmin() return idx def curve_func(x, A, B): return 8*A*(0.1875 * (x + B))**(1/3) x = np.linspace(0, 1000) y = 8*0.1875*x**(1/3) # Fit the curve initial_guess = [0.2, 16] # Initial guess for the parameters dep = 5 pos_dep = find_nearest(x[::-1], dep) params, pcov = curve_fit(curve_func, x[::-1][:pos_dep]/10, y[::-1][:pos_dep], bounds = ([0,1],[1,25])) # Extract the fitted parameters A_fit, B_fit= params # Generate the curve using the fitted parameters curve_fit_data = curve_func(x, A_fit, B_fit) # Plot the original data and the fitted curve plt.figure() plt.scatter(x[::-1][:pos_dep], y[::-1][:pos_dep], label="Data") plt.plot(x[::-1][:pos_dep], curve_fit_data[::-1][:pos_dep], color="r", label="Fit") plt.legend() plt.show()
问题排查与优化方案
1. 核心数据问题
- 自变量因变量不匹配:拟合时将输入
x除以10,但目标y仍是基于原始x计算的,直接破坏了变量间的对应关系。 - 数据区间缺失拐点:当前截取的是
x从1000到5的区间,完全避开了拐点集中的x趋近于0的区域,拟合算法无法学习到拐点特性。
2. 初始值设置偏差
对比原始数据生成公式与拟合方程,可推导真实参数(当B=0时):
$$A = 0.1875^{2/3} \approx 0.325$$
初始猜测值[0.2,16]与真实值偏差较大,容易导致拟合陷入局部最优。
3. 优化后的代码示例
import matplotlib.pyplot as plt import numpy as np from scipy.optimize import curve_fit plt.close("all") def curve_func(x, A, B): return 8*A*(0.1875 * (x + B))**(1/3) # 生成覆盖拐点区域的全区间数据,增加采样点提升精度 x = np.linspace(0, 1000, 200) y_true = 8*0.1875*x**(1/3) # 添加模拟噪声贴近真实场景 y = y_true + np.random.normal(0, 0.05, size=len(x)) # 基于真实参数设置初始值 initial_guess = [0.3, 1] # 保持原参数约束,增加迭代次数确保收敛 params, pcov = curve_fit(curve_func, x, y, bounds=([0, 1], [1, 25]), maxfev=10000) A_fit, B_fit = params print(f"拟合参数:A={A_fit:.4f}, B={B_fit:.4f}") curve_fit_data = curve_func(x, A_fit, B_fit) plt.figure() plt.scatter(x, y, s=5, label="Data") plt.plot(x, curve_fit_data, color="r", label="Fit") plt.plot(x, y_true, color="g", linestyle="--", label="True Curve") plt.xlabel("x") plt.ylabel("y") plt.legend() plt.show()
4. 额外优化建议
- 添加权重:若拐点区域数据更重要,可给
x较小的区域设置更高权重(如weights = 1/(x+1)),通过curve_fit的sigma参数传入。 - 调整精度参数:指定
method='trf'并设置ftol=1e-8、xtol=1e-8,提升拟合精度。 - 数据归一化:对
x进行归一化处理(如除以最大值),避免数值范围过大导致拟合不稳定。
内容的提问来源于stack exchange,提问作者nitika
相关产品推荐
相关产品推荐

