基于Python的非线性最小二乘拟合报错及参数求解问题
问题:非线性最小二乘拟合报错及参数求解
问题描述
尝试基于二维数据求解最优参数,参考相关内容修改代码后运行报错:
TypeError: only size-1 arrays can be converted to Python scalars
自定义方程无法支持NumPy数组输入,需解决该问题并推导参数值。
原始代码
from scipy.optimize import curve_fit import numpy as np import matplotlib.pyplot as plt t_data = np.array([0,5,10,15,20,25,27]) y_data = np.array([1771,8109,22571,30008,40862,56684,59101]) def func_nl_lsq(t, *args): K, A, L, b = args return math.exp(L*x)/((K+A*(b**x))) popt, pcov = curve_fit(func_nl_lsq, t_data, y_data, p0=[1, 1, 1, 1]) plt.plot(t_data, y_data, 'o') plt.plot(t_data, func_nl_lsq(t_data, *popt), '-') plt.show() print(popt[0], popt[1], popt[2], popt[3])
拟合公式与观测数据
拟合目标为经典增长模型,对应观测数据如下:
| 年份(x_data) | 数值(y_data) |
|---|---|
| 1960 | 1771 |
| 1965 | 8109 |
| 1970 | 22571 |
| 1975 | 30008 |
| 1980 | 40862 |
| 1985 | 56684 |
| 1987 | 59101 |
问题分析与解决
错误原因
- 函数类型不匹配:使用Python标准库
math.exp,该函数仅支持标量输入,无法处理NumPy数组,触发类型错误。 - 变量名错误:函数参数定义为
t,但公式中使用未定义的x,导致逻辑混乱。 - 参数传递冗余:
curve_fit要求自定义函数参数顺序为「自变量+待拟合参数」,*args写法易引发参数解析问题。
修正后的代码
from scipy.optimize import curve_fit import numpy as np import matplotlib.pyplot as plt # 时间差数据(以1960年为基准) t_data = np.array([0,5,10,15,20,25,27]) y_data = np.array([1771,8109,22571,30008,40862,56684,59101]) # 修正后的拟合函数:支持数组运算,变量名统一,参数格式符合curve_fit要求 def func_nl_lsq(t, K, A, L, b): return np.exp(L * t) / (K + A * (b ** t)) # 初始参数:根据数据增长趋势设置,提升拟合成功率 p0 = [1000, 500, 0.1, 1.2] popt, pcov = curve_fit(func_nl_lsq, t_data, y_data, p0=p0) # 可视化拟合结果 plt.plot(t_data, y_data, 'o', label='观测数据') plt.plot(t_data, func_nl_lsq(t_data, *popt), '-', label='拟合曲线') plt.legend() plt.xlabel('年份差(以1960为基准)') plt.ylabel('数值') plt.show() # 输出最优参数 print(f"拟合参数:K={popt[0]:.2f}, A={popt[1]:.2f}, L={popt[2]:.4f}, b={popt[3]:.4f}")
关键修正说明
- 替换
math.exp为np.exp:利用NumPy函数的向量化特性,支持数组输入运算。 - 统一变量名:将公式中的
x改为函数参数t,保证变量逻辑一致。 - 调整参数格式:直接将待拟合参数(K,A,L,b)作为显式参数,贴合
curve_fit的参数解析规则。 - 优化初始参数:根据数据量级和增长趋势设置更合理的初始值,避免拟合发散。
额外优化建议
如果拟合出现发散,可通过bounds参数限制参数取值范围,例如:
# 约束参数为正数,b在1~2之间(符合增长模型逻辑) popt, pcov = curve_fit(func_nl_lsq, t_data, y_data, p0=p0, bounds=(0, [np.inf, np.inf, 1, 2]))
内容的提问来源于stack exchange,提问作者Todd
相关产品推荐
相关产品推荐

