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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 03:27:53